From 1361ba711c2a9f81b6cb826a540b2e8511b85d0c Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Mon, 14 Sep 2026 12:43:54 +0000 Subject: [PATCH 01/39] feat(pytorch): support GLM-5.3 Flash with shared LMDeploy components Reuse MLA/MoE loaders, FA3, mHC, sparse Top-K and compact DeepGEMM. Add KDA state adaptation and multimodal GLM configuration/processing. Preserve existing defaults and the positional MoE prefix argument. --- lmdeploy/archs.py | 1 + lmdeploy/hf_configs/__init__.py | 5 +- .../hf_configs/configuration_glm5_next.py | 402 ++++ lmdeploy/pytorch/backends/attention.py | 12 + lmdeploy/pytorch/backends/base.py | 1 + .../backends/cuda/attention/default.py | 1 + .../pytorch/backends/cuda/attention/mla.py | 47 + .../cuda/attention/tilelang_sparse_mla.py | 159 ++ lmdeploy/pytorch/backends/cuda/kda.py | 179 ++ lmdeploy/pytorch/backends/cuda/kpool.py | 213 ++ .../pytorch/backends/cuda/moe/blocked_fp8.py | 96 +- lmdeploy/pytorch/backends/cuda/op_backend.py | 3 + lmdeploy/pytorch/backends/kda.py | 39 + lmdeploy/pytorch/backends/moe.py | 5 +- lmdeploy/pytorch/config.py | 10 + lmdeploy/pytorch/configurations/glm5_next.py | 311 +++ lmdeploy/pytorch/consts.py | 8 + lmdeploy/pytorch/engine/executor/base.py | 6 + .../pytorch/engine/executor/base_worker.py | 4 + .../pytorch/engine/executor/ray_executor.py | 15 +- lmdeploy/pytorch/kernels/cuda/activation.py | 22 +- lmdeploy/pytorch/kernels/cuda/moe/ep.py | 11 +- lmdeploy/pytorch/kernels/cuda/moe/ep_fp8.py | 59 +- .../kernels/cuda/moe/route_noaux_tc.py | 5 +- .../kernels/cuda/sparse_mla_tilelang.py | 234 +++ lmdeploy/pytorch/models/deepseek_v2.py | 47 +- lmdeploy/pytorch/models/deepseek_v32.py | 12 +- lmdeploy/pytorch/models/glm4_1v.py | 7 +- lmdeploy/pytorch/models/glm5_next.py | 1856 +++++++++++++++++ lmdeploy/pytorch/models/module_map.py | 6 + lmdeploy/pytorch/nn/__init__.py | 5 +- lmdeploy/pytorch/nn/attention.py | 30 + lmdeploy/pytorch/nn/kda.py | 48 + lmdeploy/pytorch/nn/kpool.py | 897 ++++++++ lmdeploy/pytorch/nn/linear/default.py | 7 +- lmdeploy/pytorch/nn/moe/__init__.py | 7 + lmdeploy/pytorch/nn/moe/blocked_fp8.py | 10 +- lmdeploy/pytorch/nn/norm.py | 39 + lmdeploy/pytorch/nn/rotary_embedding.py | 24 + .../pytorch/third_party/deep_gemm/__init__.py | 31 +- lmdeploy/serve/processors/multimodal.py | 28 + lmdeploy/vl/media/video.py | 7 +- lmdeploy/vl/media/video_loader.py | 185 +- lmdeploy/vl/model/builder.py | 1 + lmdeploy/vl/model/glm5_next.py | 295 +++ requirements/runtime_cuda.txt | 2 +- 46 files changed, 5303 insertions(+), 89 deletions(-) create mode 100644 lmdeploy/hf_configs/configuration_glm5_next.py create mode 100644 lmdeploy/pytorch/backends/cuda/attention/tilelang_sparse_mla.py create mode 100644 lmdeploy/pytorch/backends/cuda/kda.py create mode 100644 lmdeploy/pytorch/backends/cuda/kpool.py create mode 100644 lmdeploy/pytorch/backends/kda.py create mode 100644 lmdeploy/pytorch/configurations/glm5_next.py create mode 100644 lmdeploy/pytorch/kernels/cuda/sparse_mla_tilelang.py create mode 100644 lmdeploy/pytorch/models/glm5_next.py create mode 100644 lmdeploy/pytorch/nn/kda.py create mode 100644 lmdeploy/pytorch/nn/kpool.py create mode 100644 lmdeploy/vl/model/glm5_next.py diff --git a/lmdeploy/archs.py b/lmdeploy/archs.py index f9dc47c85a..ae61ef62ef 100644 --- a/lmdeploy/archs.py +++ b/lmdeploy/archs.py @@ -106,6 +106,7 @@ def check_vl_llm(backend: str, config: dict) -> bool: 'Gemma3ForConditionalGeneration', 'Llama4ForConditionalGeneration', 'InternVLForConditionalGeneration', 'InternS1ForConditionalGeneration', 'InternS1ProForConditionalGeneration', 'InternS1_1_ForConditionalGeneration', 'Glm4vForConditionalGeneration', + 'Glm5NextForConditionalGeneration', 'InternS2MobiusForConditionalGeneration', 'InternS2MobiusForCausalLM', 'InternS2PreviewForConditionalGeneration', 'InternS2PreviewForCausalLM', 'KimiK25ForConditionalGeneration', 'Kimi_K25ForConditionalGeneration', diff --git a/lmdeploy/hf_configs/__init__.py b/lmdeploy/hf_configs/__init__.py index 2a4386a5e2..db92e9c71f 100644 --- a/lmdeploy/hf_configs/__init__.py +++ b/lmdeploy/hf_configs/__init__.py @@ -11,7 +11,10 @@ @lru_cache def register_config(model_type: str): """Register an LMDeploy-owned Transformers config when available.""" - if model_type == 'kimi_k2': + if model_type == 'glm5_next': + from .configuration_glm5_next import Glm5NextConfig + AutoConfig.register(Glm5NextConfig.model_type, Glm5NextConfig) + elif model_type == 'kimi_k2': # Standalone Kimi EAGLE checkpoints do not provide an auto_map. from .configuration_kimi_k2 import KimiK2Config AutoConfig.register(KimiK2Config.model_type, KimiK2Config) diff --git a/lmdeploy/hf_configs/configuration_glm5_next.py b/lmdeploy/hf_configs/configuration_glm5_next.py new file mode 100644 index 0000000000..db35095cb7 --- /dev/null +++ b/lmdeploy/hf_configs/configuration_glm5_next.py @@ -0,0 +1,402 @@ +# Copyright (c) OpenMMLab. All rights reserved. +"""Local Transformers configuration for GLM-5.3-Flash. + +Transformers releases that predate GLM-5.3 do not know the ``glm5_next`` +model type. Keeping the configuration local lets LMDeploy read official +checkpoints without executing remote modeling code. +""" + +from __future__ import annotations + +from typing import Any + +from transformers.configuration_utils import PretrainedConfig +from transformers.models.glm_ocr.configuration_glm_ocr import GlmOcrVisionConfig + +_GLM5_NEXT_TOP_LEVEL_CONFIG_KEYS = ( + 'architectures', + 'vocab_size', + 'hidden_size', + 'head_dim', + 'intermediate_size', + 'moe_intermediate_size', + 'num_hidden_layers', + 'num_attention_heads', + 'num_key_value_heads', + 'hidden_act', + 'max_position_embeddings', + 'rms_norm_eps', + 'use_cache', + 'pad_token_id', + 'bos_token_id', + 'eos_token_id', + 'rope_theta', + 'rope_scaling', + 'rope_parameters', + 'partial_rotary_factor', + 'tie_word_embeddings', + 'attention_bias', + 'attention_dropout', + 'n_routed_experts', + 'num_experts_per_tok', + 'n_shared_experts', + 'n_group', + 'topk_group', + 'norm_topk_prob', + 'routed_scaling_factor', + 'scoring_func', + 'topk_method', + 'first_k_dense_replace', + 'moe_layer_freq', + 'moe_router_dtype', + 'output_router_logits', + 'router_aux_loss_coef', + 'q_lora_rank', + 'kv_lora_rank', + 'qk_head_dim', + 'qk_nope_head_dim', + 'qk_rope_head_dim', + 'v_head_dim', + 'mla_use_nope', + 'swiglu_limit', + 'mhc', + 'hc_mult', + 'hc_sinkhorn_iters', + 'hc_eps', + 'num_nextn_predict_layers', + 'linear_attn_config', + 'linear_head_dim', + 'linear_num_heads', + 'linear_conv_kernel_dim', + 'linear_lower_bound', + 'gate_lower_bound', + 'index_head_dim', + 'index_topk', + 'index_kpool', + 'index_kpool_always_select_tail', + 'index_kpool_compress', + 'index_n_heads', + 'index_topk_freq', + 'index_topk_pattern', + 'index_skip_topk_offset', + 'index_share_for_mtp_iteration', + 'indexer_rope_interleave', + 'indexer_types', + 'layer_types', + 'mlp_layer_types', + 'initializer_range', + 'quantization_config', +) + +_GLM5_NEXT_OUTER_ONLY_CONFIG_KEYS = frozenset({ + 'architectures', + 'quantization_config', +}) + + +class Glm5NextTextConfig(PretrainedConfig): + """Configuration of the GLM-5.3 text tower.""" + + model_type = 'glm5_next_text' + base_config_key = 'text_config' + keys_to_ignore_at_inference = ['past_key_values'] + + def __init__(self, + vocab_size: int = 154880, + hidden_size: int = 4096, + head_dim: int | None = 0, + intermediate_size: int = 12288, + moe_intermediate_size: int = 2048, + num_hidden_layers: int = 45, + num_attention_heads: int = 64, + num_key_value_heads: int | None = 64, + hidden_act: str = 'silu', + max_position_embeddings: int = 1048576, + rms_norm_eps: float = 1e-5, + use_cache: bool = True, + pad_token_id: int | None = None, + bos_token_id: int | None = None, + eos_token_id: int | list[int] | None = None, + rope_theta: float = 800000.0, + rope_scaling: dict[str, Any] | None = None, + rope_parameters: dict[str, Any] | None = None, + partial_rotary_factor: float = 1.0, + tie_word_embeddings: bool = False, + attention_bias: bool = False, + attention_dropout: float = 0.0, + n_routed_experts: int | None = 288, + num_experts_per_tok: int = 8, + n_shared_experts: int | None = 1, + n_group: int = 1, + topk_group: int = 1, + norm_topk_prob: bool = True, + routed_scaling_factor: float = 2.5, + scoring_func: str = 'sigmoid', + topk_method: str = 'noaux_tc', + first_k_dense_replace: int = 3, + moe_layer_freq: int | None = 1, + moe_router_dtype: str = 'float32', + output_router_logits: bool = False, + router_aux_loss_coef: float = 0.001, + q_lora_rank: int | None = 1536, + kv_lora_rank: int = 512, + qk_head_dim: int | None = 256, + qk_nope_head_dim: int = 256, + qk_rope_head_dim: int = 0, + v_head_dim: int = 256, + mla_use_nope: bool = True, + swiglu_limit: float | None = 10.0, + mhc: bool = True, + hc_mult: int = 4, + hc_sinkhorn_iters: int = 20, + hc_eps: float = 1e-6, + num_nextn_predict_layers: int = 1, + linear_attn_config: dict[str, Any] | None = None, + linear_head_dim: int = 128, + linear_num_heads: int = 64, + linear_conv_kernel_dim: int = 4, + linear_lower_bound: float | None = None, + gate_lower_bound: float | None = -5.0, + index_head_dim: int | None = 128, + index_topk: int | None = 2048, + index_kpool: int = 4, + index_kpool_always_select_tail: bool = True, + index_kpool_compress: bool = True, + index_n_heads: int | None = 32, + index_topk_freq: int = 1, + index_topk_pattern: str | None = None, + index_skip_topk_offset: int | None = None, + index_share_for_mtp_iteration: bool = True, + indexer_rope_interleave: bool = True, + indexer_types: list[str] | None = None, + layer_types: list[str] | None = None, + mlp_layer_types: list[str] | None = None, + initializer_range: float = 0.02, + **kwargs): + if rope_scaling is None and rope_parameters is not None: + rope_scaling = rope_parameters + if rope_parameters is not None: + rope_theta = rope_parameters.get('rope_theta', rope_theta) + partial_rotary_factor = rope_parameters.get( + 'partial_rotary_factor', partial_rotary_factor) + + if num_key_value_heads is None: + num_key_value_heads = num_attention_heads + if qk_head_dim is None: + qk_head_dim = qk_nope_head_dim + qk_rope_head_dim + # Some released configs serialize this optional field as null. The + # GLM-5.3 layout uses MoE on every layer after the dense prefix. + if moe_layer_freq is None: + moe_layer_freq = 1 + + self.vocab_size = vocab_size + self.hidden_size = hidden_size + self.head_dim = head_dim + self.intermediate_size = intermediate_size + self.moe_intermediate_size = moe_intermediate_size + self.num_hidden_layers = num_hidden_layers + self.num_attention_heads = num_attention_heads + self.num_key_value_heads = num_key_value_heads + self.hidden_act = hidden_act + self.max_position_embeddings = max_position_embeddings + self.rms_norm_eps = rms_norm_eps + self.use_cache = use_cache + self.rope_theta = rope_theta + self.rope_scaling = rope_scaling + self.rope_parameters = rope_parameters + self.partial_rotary_factor = partial_rotary_factor + self.attention_bias = attention_bias + self.attention_dropout = attention_dropout + self.n_routed_experts = n_routed_experts + self.num_experts_per_tok = num_experts_per_tok + self.n_shared_experts = n_shared_experts + self.n_group = n_group + self.topk_group = topk_group + self.norm_topk_prob = norm_topk_prob + self.routed_scaling_factor = routed_scaling_factor + self.scoring_func = scoring_func + self.topk_method = topk_method + self.first_k_dense_replace = first_k_dense_replace + self.moe_layer_freq = moe_layer_freq + self.moe_router_dtype = moe_router_dtype + self.output_router_logits = output_router_logits + self.router_aux_loss_coef = router_aux_loss_coef + self.q_lora_rank = q_lora_rank + self.kv_lora_rank = kv_lora_rank + self.qk_head_dim = qk_head_dim + self.qk_nope_head_dim = qk_nope_head_dim + self.qk_rope_head_dim = qk_rope_head_dim + self.v_head_dim = v_head_dim + self.mla_use_nope = mla_use_nope + self.swiglu_limit = swiglu_limit + self.mhc = mhc + self.hc_mult = hc_mult + self.hc_sinkhorn_iters = hc_sinkhorn_iters + self.hc_eps = hc_eps + self.num_nextn_predict_layers = num_nextn_predict_layers + self.index_head_dim = index_head_dim + self.index_topk = index_topk + self.index_kpool = index_kpool + self.index_kpool_always_select_tail = index_kpool_always_select_tail + self.index_kpool_compress = index_kpool_compress + self.index_n_heads = index_n_heads + self.index_topk_freq = index_topk_freq + self.index_topk_pattern = index_topk_pattern + self.index_skip_topk_offset = index_skip_topk_offset + self.index_share_for_mtp_iteration = index_share_for_mtp_iteration + self.indexer_rope_interleave = indexer_rope_interleave + self.indexer_types = indexer_types + self.layer_types = layer_types + self.mlp_layer_types = mlp_layer_types + self.initializer_range = initializer_range + self.linear_lower_bound = linear_lower_bound + self.gate_lower_bound = (gate_lower_bound + if gate_lower_bound is not None else + linear_lower_bound) + + if linear_attn_config is None: + if layer_types is None: + kda_layers = [ + layer_idx for layer_idx in range(num_hidden_layers) + if layer_idx % 4 != 3 + ] + else: + kda_layers = [ + layer_idx for layer_idx, layer_type in enumerate(layer_types) + if layer_type == 'linear_attention' + ] + kda_layer_set = set(kda_layers) + linear_attn_config = { + 'full_attn_layers': [ + layer_idx for layer_idx in range(num_hidden_layers) + if layer_idx not in kda_layer_set + ], + 'head_dim': linear_head_dim, + 'kda_layers': kda_layers, + 'num_heads': linear_num_heads, + 'short_conv_kernel_size': linear_conv_kernel_dim, + 'gate_lower_bound': self.gate_lower_bound, + } + else: + linear_attn_config = dict(linear_attn_config) + + linear_head_dim = linear_attn_config.get('head_dim', + linear_head_dim) + linear_num_heads = linear_attn_config.get('num_heads', + linear_num_heads) + linear_conv_kernel_dim = linear_attn_config.get( + 'short_conv_kernel_size', linear_conv_kernel_dim) + if gate_lower_bound is None: + gate_lower_bound = linear_attn_config.get('gate_lower_bound', + linear_lower_bound) + linear_attn_config.setdefault('head_dim', linear_head_dim) + linear_attn_config.setdefault('num_heads', linear_num_heads) + linear_attn_config.setdefault('short_conv_kernel_size', + linear_conv_kernel_dim) + linear_attn_config.setdefault('gate_lower_bound', gate_lower_bound) + + self.linear_head_dim = linear_head_dim + self.linear_num_heads = linear_num_heads + self.linear_conv_kernel_dim = linear_conv_kernel_dim + self.gate_lower_bound = gate_lower_bound + self.linear_attn_config = linear_attn_config + + super().__init__(pad_token_id=pad_token_id, + bos_token_id=bos_token_id, + eos_token_id=eos_token_id, + tie_word_embeddings=tie_word_embeddings, + **kwargs) + if rope_parameters is not None or rope_scaling is not None: + self.rope_parameters = rope_parameters or rope_scaling + + def is_kda_layer(self, layer_idx: int) -> bool: + """Return whether ``layer_idx`` uses KDA instead of sparse MLA.""" + return (self.linear_attn_config is not None + and layer_idx in self.linear_attn_config['kda_layers']) + + @property + def linear_layer_ids(self) -> list[int]: + return [ + layer_idx for layer_idx in range(self.num_hidden_layers) + if self.is_kda_layer(layer_idx) + ] + + @property + def full_attention_layer_ids(self) -> list[int]: + return [ + layer_idx for layer_idx in range(self.num_hidden_layers) + if not self.is_kda_layer(layer_idx) + ] + + @property + def nextn_layer_ids(self) -> list[int]: + return [ + self.num_hidden_layers + layer_idx + for layer_idx in range(self.num_nextn_predict_layers or 0) + ] + + +class Glm5NextVisionConfig(GlmOcrVisionConfig): + """GLM-OCR vision configuration with clamped SwiGLU.""" + + model_type = 'glm5_next_vision' + + def __init__(self, swiglu_limit: float = 10.0, **kwargs): + super().__init__(**kwargs) + self.swiglu_limit = swiglu_limit + + +class Glm5NextConfig(PretrainedConfig): + """Top-level GLM-5.3 multimodal configuration.""" + + model_type = 'glm5_next' + sub_configs = { + 'vision_config': Glm5NextVisionConfig, + 'text_config': Glm5NextTextConfig, + } + keys_to_ignore_at_inference = ['past_key_values'] + + def __init__(self, + text_config: dict[str, Any] | PretrainedConfig | None = None, + vision_config: dict[str, Any] | PretrainedConfig | None = None, + image_token_id: int = 154854, + video_token_id: int = 154855, + image_start_token_id: int = 154830, + image_end_token_id: int = 154831, + video_start_token_id: int = 154832, + video_end_token_id: int = 154833, + **kwargs): + top_level_text_config = { + key: kwargs[key] + for key in _GLM5_NEXT_TOP_LEVEL_CONFIG_KEYS + if key in kwargs and key not in _GLM5_NEXT_OUTER_ONLY_CONFIG_KEYS + } + if isinstance(text_config, dict): + text_config = {**top_level_text_config, **text_config} + text_config = Glm5NextTextConfig(**text_config) + elif text_config is None: + text_config = Glm5NextTextConfig(**top_level_text_config) + self.text_config = text_config + + if isinstance(vision_config, dict): + vision_config = Glm5NextVisionConfig(**vision_config) + self.vision_config = vision_config + self.image_token_id = image_token_id + self.video_token_id = video_token_id + self.image_start_token_id = image_start_token_id + self.image_end_token_id = image_end_token_id + self.video_start_token_id = video_start_token_id + self.video_end_token_id = video_end_token_id + + if getattr(self.text_config, 'quantization_config', None) is not None: + self.quantization_config = self.text_config.quantization_config + + super().__init__(**kwargs) + + # Existing LMDeploy configuration builders still read language fields + # from the outer config. Mirror language fields without overwriting + # top-level model-selection or quantization metadata. + for key in _GLM5_NEXT_TOP_LEVEL_CONFIG_KEYS: + if (key not in _GLM5_NEXT_OUTER_ONLY_CONFIG_KEYS + and hasattr(self.text_config, key)): + setattr(self, key, getattr(self.text_config, key)) diff --git a/lmdeploy/pytorch/backends/attention.py b/lmdeploy/pytorch/backends/attention.py index 37804fbe15..f5994afb31 100644 --- a/lmdeploy/pytorch/backends/attention.py +++ b/lmdeploy/pytorch/backends/attention.py @@ -160,6 +160,18 @@ def make_alibi_slopes(head_start: int, head_end: int, num_heads: int, alibi_scal def set_alibi_slopes(self, slopes: torch.Tensor): self.alibi_slopes = slopes + def fill_and_flatten_latent_kv_cache( + self, + key: torch.Tensor, + k_cache: torch.Tensor, + attn_metadata: T, + out_dtype: torch.dtype = None, + k_scales_zeros: torch.Tensor = None, + v_scales_zeros: torch.Tensor = None, + ) -> torch.Tensor: + """Append latent KV and return the request-major prefill view.""" + raise NotImplementedError(f'{type(self).__name__} does not support latent KV cache flattening.') + @abstractmethod def forward( self, diff --git a/lmdeploy/pytorch/backends/base.py b/lmdeploy/pytorch/backends/base.py index a6ffaf1c22..4f7eaadfa0 100644 --- a/lmdeploy/pytorch/backends/base.py +++ b/lmdeploy/pytorch/backends/base.py @@ -51,6 +51,7 @@ class OpType(Enum): # Gated Delta CausalConv1d = auto() GatedDeltaRule = auto() + Kda = auto() class OpsBackend(ABC): diff --git a/lmdeploy/pytorch/backends/cuda/attention/default.py b/lmdeploy/pytorch/backends/cuda/attention/default.py index a627cb45e0..416141d399 100644 --- a/lmdeploy/pytorch/backends/cuda/attention/default.py +++ b/lmdeploy/pytorch/backends/cuda/attention/default.py @@ -73,6 +73,7 @@ def build_triton_attention_metadata(attn_meta_cls, step_context, cu_seqlens_q=sequence_metadata.cu_seqlens_q, cu_seqlens_k=sequence_metadata.cu_seqlens_k, max_kv_seqlen=sequence_metadata.max_kv_seqlen, + max_q_seqlen=step_context.max_q_seqlen, ) diff --git a/lmdeploy/pytorch/backends/cuda/attention/mla.py b/lmdeploy/pytorch/backends/cuda/attention/mla.py index 9178209c68..b5b683010a 100644 --- a/lmdeploy/pytorch/backends/cuda/attention/mla.py +++ b/lmdeploy/pytorch/backends/cuda/attention/mla.py @@ -510,6 +510,53 @@ def _fill_kv_cache_impl(self, block_offsets=block_offsets, ) + def fill_and_flatten_latent_kv_cache( + self, + key: torch.Tensor, + k_cache: torch.Tensor, + attn_metadata: TritonAttentionMetadata, + out_dtype: torch.dtype = None, + k_scales_zeros: torch.Tensor = None, + v_scales_zeros: torch.Tensor = None, + ) -> torch.Tensor: + """Append latent KV and flatten the complete prefill cache. + + ``key`` contains only the current tokens. The paged cache may already + contain a prefix; ``attn_metadata`` determines where current tokens are + appended and how every request is flattened into ``shd`` layout. + """ + if attn_metadata.is_decoding: + raise RuntimeError('Latent KV cache flattening is only supported during prefill.') + if out_dtype is None: + out_dtype = key.dtype + + # MLA carries its latent value payload in the leading K dimensions. + # Aliased views keep the existing fill/flatten kernels on their + # shared-KV fast path without a duplicate value store. + value = key[..., :self.v_head_size] + v_cache = k_cache[..., :self.v_head_size] + max_q_seqlen = self._get_max_q_seqlen(key, attn_metadata) + self._fill_kv_cache_impl( + key, + value, + k_cache, + v_cache, + attn_metadata, + max_q_seqlen, + k_scales_zeros=k_scales_zeros, + v_scales_zeros=v_scales_zeros, + ) + flatten_k, _ = self._flatten_prefill_kv_cache( + k_cache, + v_cache, + attn_metadata, + out_dtype=out_dtype, + kv_layout='shd', + k_scales_zeros=k_scales_zeros, + v_scales_zeros=v_scales_zeros, + ) + return flatten_k + def _forward_decoding( self, query: torch.Tensor, diff --git a/lmdeploy/pytorch/backends/cuda/attention/tilelang_sparse_mla.py b/lmdeploy/pytorch/backends/cuda/attention/tilelang_sparse_mla.py new file mode 100644 index 0000000000..f373067c58 --- /dev/null +++ b/lmdeploy/pytorch/backends/cuda/attention/tilelang_sparse_mla.py @@ -0,0 +1,159 @@ +# Copyright (c) OpenMMLab. All rights reserved. +"""TileLang sparse-MLA backend with reusable logical-index mapping.""" + +from __future__ import annotations + +from typing import Any + +import torch + +from lmdeploy.pytorch.kernels.cuda.sparse_mla_tilelang import ( + sparse_mla_bf16_fwd, +) + +from .sparse_mla import FlashMLAIndexMapper + + +class TilelangSparseMLADecode: + """Run sparse MLA from chronological or model-provided logical indices.""" + + def __init__(self, index_topk: int, index_kpool: int = 1): + if index_topk <= 0: + raise ValueError(f'index_topk must be positive, got {index_topk}') + if index_kpool <= 0: + raise ValueError(f'index_kpool must be positive, got {index_kpool}') + self.index_topk = index_topk + self.index_kpool = index_kpool + tail_width = index_kpool - 1 + self.kernel_topk = ((index_topk + tail_width + 63) // 64) * 64 + self.index_mapper = FlashMLAIndexMapper.build() + + def _pad_logical_indices(self, indices: torch.Tensor) -> torch.Tensor: + """Pad model indices with -1 to TileLang's 64-entry tile width.""" + if indices.ndim != 2: + raise ValueError( + 'Sparse MLA logical indices must have shape [tokens, topk], ' + f'got {tuple(indices.shape)}.') + if indices.size(1) > self.kernel_topk: + raise ValueError( + f'Logical index width {indices.size(1)} exceeds the ' + f'kernel width {self.kernel_topk}.') + if indices.dtype != torch.int32: + indices = indices.to(torch.int32) + return torch.nn.functional.pad( + indices, (0, self.kernel_topk - indices.size(1)), value=-1) + + def _build_physical_indices(self, query: torch.Tensor, + k_cache: torch.Tensor, + attn_metadata: Any) -> torch.Tensor: + kv_seqlens = attn_metadata.kv_seqlens + block_offsets = attn_metadata.block_offsets + batch_size = kv_seqlens.numel() + if query.size(0) != batch_size: + raise NotImplementedError( + 'TileLang sparse MLA currently requires one decode token per ' + f'request, got {query.size(0)} tokens for {batch_size} requests.') + if int(attn_metadata.max_kv_seqlen) > self.index_topk: + raise NotImplementedError( + 'Contexts longer than index_topk require the model KPool indexer.') + + logical = torch.arange(self.kernel_topk, + dtype=torch.int32, + device=query.device) + logical = logical.unsqueeze(0).expand(batch_size, -1).clone() + logical.masked_fill_(logical >= kv_seqlens[:, None], -1) + return self.index_mapper.map_paged_decode( + logical, + block_offsets, + max_q_seqlen=1, + block_size=k_cache.size(1), + ) + + def _map_decode_indices(self, logical_indices: torch.Tensor, + k_cache: torch.Tensor, + attn_metadata: Any) -> torch.Tensor: + logical_indices = self._pad_logical_indices(logical_indices) + batch_size = attn_metadata.kv_seqlens.numel() + if logical_indices.size(0) % batch_size: + raise ValueError( + 'Decode logical-index rows must be divisible by batch size, ' + f'got rows={logical_indices.size(0)}, batch={batch_size}.') + max_q_seqlen = logical_indices.size(0) // batch_size + return self.index_mapper.map_paged_decode( + logical_indices, + attn_metadata.block_offsets, + max_q_seqlen=max_q_seqlen, + block_size=k_cache.size(1), + ).flatten(0, 1)[:, None] + + def _map_prefill_indices(self, logical_indices: torch.Tensor, + attn_metadata: Any) -> torch.Tensor: + logical_indices = self._pad_logical_indices(logical_indices) + return self.index_mapper.map_flat_prefill( + logical_indices, + attn_metadata.q_seqlens, + attn_metadata.cu_seqlens_k, + ) + + def forward(self, + query: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + k_cache: torch.Tensor, + v_cache: torch.Tensor, + attn_metadata: Any, + scale: float, + cache_writer: Any, + k_scales_zeros: torch.Tensor | None = None, + v_scales_zeros: torch.Tensor | None = None, + logical_indices: torch.Tensor | None = None) -> torch.Tensor: + """Append latent KV, then run BF16 sparse MLA decode.""" + if k_cache.dtype != torch.bfloat16: + raise TypeError('TileLang sparse MLA requires a BF16 KV cache.') + cache_writer._lazy_init(query.device) + cache_impl = cache_writer.impl + max_q_seqlen = cache_impl._get_max_q_seqlen(query, attn_metadata) + cache_impl._fill_kv_cache_impl( + key, + value, + k_cache, + v_cache, + attn_metadata, + max_q_seqlen, + k_scales_zeros=k_scales_zeros, + v_scales_zeros=v_scales_zeros, + ) + if logical_indices is None: + indices = self._build_physical_indices( + query, k_cache, attn_metadata).flatten(0, 1)[:, None] + else: + indices = self._map_decode_indices( + logical_indices, k_cache, attn_metadata) + flat_cache = k_cache.flatten(0, 1) + return sparse_mla_bf16_fwd(query, flat_cache, indices, scale) + + def forward_prefill( + self, + query: torch.Tensor, + key: torch.Tensor, + k_cache: torch.Tensor, + attn_metadata: Any, + scale: float, + cache_writer: Any, + logical_indices: torch.Tensor, + k_scales_zeros: torch.Tensor | None = None, + v_scales_zeros: torch.Tensor | None = None, + ) -> torch.Tensor: + """Append/flatten latent KV, then run sparse MLA prefill.""" + if k_cache.dtype != torch.bfloat16: + raise TypeError('TileLang sparse MLA requires a BF16 KV cache.') + flat_cache = cache_writer.fill_and_flatten_latent_kv_cache( + key, + k_cache, + attn_metadata, + out_dtype=query.dtype, + k_scales_zeros=k_scales_zeros, + v_scales_zeros=v_scales_zeros, + ) + indices = self._map_prefill_indices(logical_indices, attn_metadata) + return sparse_mla_bf16_fwd(query, flat_cache, indices, scale) diff --git a/lmdeploy/pytorch/backends/cuda/kda.py b/lmdeploy/pytorch/backends/cuda/kda.py new file mode 100644 index 0000000000..9753c0d949 --- /dev/null +++ b/lmdeploy/pytorch/backends/cuda/kda.py @@ -0,0 +1,179 @@ +# Copyright (c) OpenMMLab. All rights reserved. +"""CUDA KDA backend composed from the public FLA operators. + +KDA is distinct from LMDeploy's gated-delta rule, but its CUDA implementation +does not need copied GLM kernels. This adapter owns only LMDeploy cache/state +semantics and delegates convolution and recurrence to FLA. +""" + +from typing import Any + +import torch + +from lmdeploy.pytorch.backends.kda import KdaBuilder, KdaImpl + + +def _select_state(state: torch.Tensor, metadata: Any) -> torch.Tensor: + selected = state.index_select(0, metadata.state_ids.long()) + clear = ~metadata.valid_state + if metadata.is_init is not None: + clear = clear | metadata.is_init + clear = clear.reshape(-1, *((1, ) * (state.ndim - 1))) + return selected.masked_fill(clear, 0) + + +def _store_state(state: torch.Tensor, value: torch.Tensor, + metadata: Any) -> None: + state_ids = metadata.state_ids.long() + valid = metadata.valid_state.reshape(-1, + *((1, ) * (state.ndim - 1))) + previous = state.index_select(0, state_ids) + stored = torch.where(valid, value.to(state.dtype), previous) + state.index_copy_(0, state_ids, stored) + + +class CudaKdaImpl(KdaImpl): + """KDA implemented by public FLA convolution/recurrence kernels.""" + + def __init__(self): + try: + from fla.modules.conv.triton.ops import ( + causal_conv1d_fwd, + causal_conv1d_update, + ) + from fla.ops.kda import chunk_kda, fused_recurrent_kda + except (ImportError, AttributeError) as exc: + raise ImportError( + 'GLM-5.3 KDA requires flash-linear-attention==0.5.2.' + ) from exc + self.causal_conv1d_fwd = causal_conv1d_fwd + self.causal_conv1d_update = causal_conv1d_update + self.chunk_kda = chunk_kda + self.fused_recurrent_kda = fused_recurrent_kda + + def _conv( + self, + mixed_qkv: torch.Tensor, + conv_weight: torch.Tensor, + conv_bias: torch.Tensor | None, + conv_state: torch.Tensor, + metadata: Any, + ) -> torch.Tensor: + selected_state = _select_state(conv_state, metadata) + if conv_weight.dim() == 3: + if conv_weight.size(1) != 1: + raise ValueError( + 'KDA depthwise convolution weight must have shape ' + '[D, 1, K].') + conv_weight = conv_weight.squeeze(1) + if metadata.is_decoding: + mixed_qkv, final_state = self.causal_conv1d_update( + x=mixed_qkv, + cache=selected_state, + weight=conv_weight, + bias=conv_bias, + activation='silu', + ) + else: + mixed_qkv, final_state = self.causal_conv1d_fwd( + x=mixed_qkv, + weight=conv_weight, + bias=conv_bias, + residual=None, + initial_state=selected_state, + output_final_state=True, + activation='silu', + cu_seqlens=metadata.cu_seqlens, + ) + _store_state(conv_state, final_state, metadata) + return mixed_qkv + + def forward( + self, + mixed_qkv: torch.Tensor, + raw_gate: torch.Tensor, + raw_beta: torch.Tensor, + conv_weight: torch.Tensor, + conv_bias: torch.Tensor | None, + a_log: torch.Tensor, + dt_bias: torch.Tensor, + conv_state: torch.Tensor, + recurrent_state: torch.Tensor, + metadata: Any, + num_heads: int, + head_dim: int, + lower_bound: float, + ) -> torch.Tensor: + if (metadata.spec_state_offsets is not None + or (getattr(metadata, 'num_spec_tokens', 0) or 0) > 0): + raise NotImplementedError( + 'GLM-5.3 KDA speculative state rollback is not implemented.') + batch_size = metadata.state_ids.numel() + if metadata.is_decoding: + query_length = mixed_qkv.size(1) // batch_size + if query_length != 1: + raise NotImplementedError( + 'GLM-5.3 KDA supports single-token autoregressive decode only.') + + mixed_qkv = self._conv(mixed_qkv, conv_weight, conv_bias, + conv_state, metadata) + q, k, v = mixed_qkv.split(num_heads * head_dim, dim=-1) + q = q.unflatten(-1, (num_heads, head_dim)).contiguous() + k = k.unflatten(-1, (num_heads, head_dim)).contiguous() + v = v.unflatten(-1, (num_heads, head_dim)).contiguous() + raw_gate = raw_gate.unflatten( + -1, (num_heads, head_dim)).contiguous() + raw_beta = raw_beta.contiguous() + selected_state = _select_state(recurrent_state, metadata) + + if metadata.is_decoding: + def decode_view(x: torch.Tensor) -> torch.Tensor: + return x.squeeze(0).unflatten( + 0, (batch_size, 1)).contiguous() + + output, final_state = self.fused_recurrent_kda( + q=decode_view(q), + k=decode_view(k), + v=decode_view(v), + g=decode_view(raw_gate), + beta=decode_view(raw_beta), + A_log=a_log, + dt_bias=dt_bias, + initial_state=selected_state, + output_final_state=True, + use_qk_l2norm_in_kernel=True, + use_gate_in_kernel=True, + use_beta_sigmoid_in_kernel=True, + lower_bound=lower_bound, + state_v_first=True, + ) + output = output.flatten(0, 1).unsqueeze(0) + else: + output, final_state = self.chunk_kda( + q=q, + k=k, + v=v, + g=raw_gate, + beta=raw_beta, + A_log=a_log, + dt_bias=dt_bias, + initial_state=selected_state, + output_final_state=True, + use_qk_l2norm_in_kernel=True, + use_gate_in_kernel=True, + use_beta_sigmoid_in_kernel=True, + safe_gate=True, + lower_bound=lower_bound, + state_v_first=True, + cu_seqlens=metadata.cu_seqlens, + ) + _store_state(recurrent_state, final_state, metadata) + return output + + +class CudaKdaBuilder(KdaBuilder): + """Build the CUDA KDA implementation.""" + + @staticmethod + def build() -> KdaImpl: + return CudaKdaImpl() diff --git a/lmdeploy/pytorch/backends/cuda/kpool.py b/lmdeploy/pytorch/backends/cuda/kpool.py new file mode 100644 index 0000000000..b2556e3c2b --- /dev/null +++ b/lmdeploy/pytorch/backends/cuda/kpool.py @@ -0,0 +1,213 @@ +# Copyright (c) OpenMMLab. All rights reserved. +"""LMDeploy CUDA adapters for pooled DSA selection.""" + +from __future__ import annotations + +import functools + +import torch +from torch import Tensor + +from lmdeploy.pytorch.kernels.cuda.sparse_index_topk import ( + is_sparse_index_topk_supported, + sparse_index_topk, +) +from lmdeploy.pytorch.nn.kpool import kpool_compress, kpool_quantize_fp8 + + +@functools.lru_cache +def _get_deep_gemm(): + try: + import deep_gemm + except ImportError as error: + raise RuntimeError( + 'GLM-5.3 KPool scoring requires DeepGEMM.') from error + required = ( + 'fp8_mqa_logits', + 'fp8_paged_mqa_logits', + 'get_paged_mqa_logits_metadata', + 'get_num_sms', + ) + missing = [name for name in required if not hasattr(deep_gemm, name)] + if missing: + raise RuntimeError( + f'GLM-5.3 KPool requires DeepGEMM APIs: {missing}.') + return deep_gemm + + +def kpool_compress_quantize_cuda( + slot_k: Tensor, + slot_score: Tensor, + ape: Tensor, + *, + mode: str, + round_scale: bool, +) -> tuple[Tensor, Tensor]: + """Compress and quantize closed pools with LMDeploy's reusable semantics.""" + pooled = kpool_compress(slot_k, slot_score, ape, mode=mode) + return kpool_quantize_fp8( + pooled, + block_size=pooled.size(-1), + round_scale=round_scale, + ) + + +def kpool_select_groups_cuda( + logits: Tensor, + group_lengths: Tensor, + *, + group_topk: int, + row_starts: Tensor | None = None, + max_group_length: int | None = None, +) -> Tensor: + """Select pooled groups with LMDeploy's shared sparse-index Top-K kernel.""" + if not is_sparse_index_topk_supported(group_topk): + raise ValueError( + 'The GLM-5.3 KPool selector only supports group_topk=512 or 2048, ' + f'got {group_topk}.') + if logits.ndim != 2 or logits.dtype != torch.float32: + raise ValueError( + 'KPool logits must be a two-dimensional float32 tensor, got ' + f'shape={tuple(logits.shape)}, dtype={logits.dtype}.') + if logits.stride(-1) != 1: + raise ValueError('KPool logits must be contiguous on the group axis.') + if group_lengths.shape != (logits.size(0), ): + raise ValueError('group_lengths must contain one value per logits row.') + if row_starts is not None and row_starts.shape != (logits.size(0), ): + raise ValueError('row_starts must contain one value per logits row.') + if max_group_length is None: + max_group_length = logits.size(1) + if max_group_length < 0 or max_group_length > logits.size(1): + raise ValueError( + 'max_group_length must be inside the logits width, got ' + f'{max_group_length} for width={logits.size(1)}.') + if max_group_length == 0: + return torch.full( + (logits.size(0), group_topk), + -1, + dtype=torch.int32, + device=logits.device, + ) + lengths = group_lengths.to( + device=logits.device, dtype=torch.int32).contiguous() + if row_starts is not None: + row_starts = row_starts.to( + device=logits.device, dtype=torch.int32).contiguous() + score_window = logits[:, :max_group_length] + if row_starts is not None: + columns = torch.arange( + max_group_length, dtype=torch.int64, device=logits.device) + gather_ids = row_starts.to(torch.int64)[:, None] + columns[None] + gather_ids = gather_ids.clamp(max=logits.size(1) - 1) + score_window = logits.gather(1, gather_ids) + q_seqlens = torch.ones( + logits.size(0), dtype=torch.int32, device=logits.device) + return sparse_index_topk( + score_window.contiguous(), + q_seqlens, + lengths.clamp(max=max_group_length), + group_topk, + fill=-1, + descending=True, + sorted=False, + ) + + +def _validate_query(query_fp8: Tensor, query_weight: Tensor) -> None: + if query_fp8.ndim != 3: + raise ValueError('query_fp8 must have shape [rows, heads, head_dim].') + if query_weight.shape != query_fp8.shape[:2]: + raise ValueError( + 'query_weight must have shape [rows, heads], got ' + f'{tuple(query_weight.shape)} for query {tuple(query_fp8.shape)}.') + if query_fp8.dtype != torch.float8_e4m3fn: + raise TypeError( + f'query_fp8 must use float8_e4m3fn, got {query_fp8.dtype}.') + if query_weight.dtype != torch.float32: + raise TypeError( + f'query_weight must use float32, got {query_weight.dtype}.') + + +def kpool_score_contiguous_cuda( + query_fp8: Tensor, + query_weight: Tensor, + pooled_key_fp8: Tensor, + pooled_key_scale: Tensor, + group_lengths: Tensor, +) -> Tensor: + """Score ragged pooled history with DeepGEMM's contiguous MQA primitive.""" + _validate_query(query_fp8, query_weight) + rows = query_fp8.size(0) + if pooled_key_fp8.ndim != 2 or pooled_key_fp8.size(1) != query_fp8.size(2): + raise ValueError( + 'pooled_key_fp8 must have shape [groups, query_head_dim].') + if pooled_key_scale.shape == (pooled_key_fp8.size(0), 1): + pooled_key_scale = pooled_key_scale.squeeze(1) + if pooled_key_scale.shape != (pooled_key_fp8.size(0), ): + raise ValueError('pooled_key_scale must contain one scale per group.') + if group_lengths.shape != (rows, ): + raise ValueError('group_lengths must contain one value per query row.') + if pooled_key_fp8.size(0) == 0: + return torch.empty( + (rows, 0), dtype=torch.float32, device=query_fp8.device) + + starts = torch.zeros(rows, dtype=torch.int32, device=query_fp8.device) + ends = group_lengths.to(device=query_fp8.device, + dtype=torch.int32).contiguous() + return _get_deep_gemm().fp8_mqa_logits( + query_fp8.contiguous(), + (pooled_key_fp8.contiguous(), pooled_key_scale.contiguous()), + query_weight.contiguous(), + starts, + ends, + clean_logits=True, + ) + + +def kpool_score_paged_cuda( + query_fp8: Tensor, + query_weight: Tensor, + packed_cache: Tensor, + group_lengths: Tensor, + pooled_block_offsets: Tensor, + page_size: int = 64, +) -> Tensor: + """Score pooled decode history with DeepGEMM's paged MQA primitive.""" + _validate_query(query_fp8, query_weight) + rows = query_fp8.size(0) + if packed_cache.dtype != torch.uint8 or packed_cache.ndim != 4: + raise ValueError( + 'packed_cache must be a uint8 [blocks, entries, 1, width] tensor.') + if packed_cache.size(1) != page_size or packed_cache.size(2) != 1: + raise ValueError( + f'packed_cache must contain [{page_size}, 1] entries per page.') + if group_lengths.shape != (rows, ): + raise ValueError('group_lengths must contain one value per query row.') + if pooled_block_offsets.ndim != 2: + raise ValueError('pooled_block_offsets must have shape [rows, pages].') + if pooled_block_offsets.size(0) == 1 and rows != 1: + pooled_block_offsets = pooled_block_offsets.expand(rows, -1) + if pooled_block_offsets.size(0) != rows: + raise ValueError( + 'pooled_block_offsets must have one row per query row.') + if pooled_block_offsets.size(1) == 0: + return torch.empty( + (rows, 0), dtype=torch.float32, device=query_fp8.device) + + deep_gemm = _get_deep_gemm() + context_lens = group_lengths.to( + device=query_fp8.device, dtype=torch.int32).contiguous().view(-1, 1) + block_table = pooled_block_offsets.to( + device=query_fp8.device, dtype=torch.int32).contiguous() + schedule = deep_gemm.get_paged_mqa_logits_metadata( + context_lens.clamp(min=1), page_size, deep_gemm.get_num_sms()) + return deep_gemm.fp8_paged_mqa_logits( + query_fp8.contiguous().unsqueeze(1), + packed_cache, + query_weight.contiguous(), + context_lens, + block_table, + schedule, + block_table.size(1) * page_size, + clean_logits=False, + ) diff --git a/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py b/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py index b33e849f39..9dd5f98d6f 100644 --- a/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py +++ b/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py @@ -28,7 +28,17 @@ logger = get_logger('lmdeploy') -class FusedMoENormal: +def _count_tokens_per_expert(topk_ids: torch.Tensor, + num_experts: int) -> torch.Tensor: + """Count routed assignments with a fixed-size, graph-safe CUDA output.""" + flat_ids = topk_ids.flatten().to(torch.int64) + counts = torch.zeros( + num_experts, dtype=torch.int32, device=topk_ids.device) + return counts.scatter_add_(0, flat_ids, torch.ones_like( + flat_ids, dtype=counts.dtype)) + + +class FusedMoENormal(FusedMoEBlockedF8Impl): def __init__( self, @@ -45,24 +55,34 @@ def __init__( num_max_dispatch_tokens_per_rank: int = 128, chunk_size: int | None = 32 * 1024, expert_alignment: int = 128, + renormalize: bool = False, + fp32_acc: bool = False, + output_scale: float = 1.0, ): + super().__init__() self.layer_index = layer_index self.top_k = top_k self.num_experts = num_experts self.block_size = block_size + self.ep_size = ep_size self.num_local_experts = num_experts // ep_size self.out_dtype = out_dtype self.fp8_dtype = fp8_dtype self.scale_fmt = scale_fmt - self.token_dispatcher = DeepEPTokenDispatcherNormal( - group=ep_group, - num_experts=num_experts, - num_local_experts=self.num_local_experts, - hidden_size=hidden_dim, - params_dtype=out_dtype, - num_max_dispatch_tokens_per_rank=num_max_dispatch_tokens_per_rank, - expert_alignment=expert_alignment, - ) + self.renormalize = renormalize + self.fp32_acc = fp32_acc + self.output_scale = output_scale + self.token_dispatcher = None + if ep_size > 1: + self.token_dispatcher = DeepEPTokenDispatcherNormal( + group=ep_group, + num_experts=num_experts, + num_local_experts=self.num_local_experts, + hidden_size=hidden_dim, + params_dtype=out_dtype, + num_max_dispatch_tokens_per_rank=num_max_dispatch_tokens_per_rank, + expert_alignment=expert_alignment, + ) def forward( self, @@ -74,7 +94,34 @@ def forward( down_weights: torch.Tensor, down_scale: torch.Tensor, expert_list: list[int] = None, + gate_up_bias: torch.Tensor = None, + down_bias: torch.Tensor = None, + act_func: Callable = None, ): + if self.token_dispatcher is None: + assert expert_list is None + assert gate_up_bias is None and down_bias is None + input_size = hidden_states.shape + hidden_states = hidden_states.flatten(0, -2) + topk_ids = topk_ids.flatten(0, -2) + topk_weights = _renormalize(topk_weights.flatten(0, -2), self.renormalize) + hs_quant, hs_scale = per_token_group_quant_fp8(hidden_states, + self.block_size, + dtype=up_weights.dtype, + scale_fmt=self.scale_fmt) + tokens_per_expert = _count_tokens_per_expert( + topk_ids, self.num_experts) + out_states = fused_moe_v3_fp8((hs_quant, hs_scale), + topk_ids, + topk_weights, (up_weights, up_scale), + (down_weights, down_scale), + tokens_per_expert, + compact_layout=True, + act_func=act_func, + fp32_acc=self.fp32_acc, + output_scale=self.output_scale) + return out_states.unflatten(0, input_size[:-1]) + hs_quant, hs_scale = per_token_group_quant_fp8(hidden_states, self.block_size, dtype=up_weights.dtype, @@ -90,6 +137,7 @@ def forward( return self.token_dispatcher.combine(out_states) def capture(self): + assert self.token_dispatcher is not None return self.token_dispatcher.buffer_normal.capture() def wait(self, event): @@ -464,9 +512,14 @@ def build(top_k: int, fp8_dtype: torch.dtype = torch.float8_e4m3fn, num_max_dispatch_tokens_per_rank: int = 128, layer_idx: int = 0, - custom_gateup_act: bool = False): + custom_gateup_act: bool = False, + fp32_acc: bool = False, + output_scale: float = 1.0, + use_deep_gemm: bool = False): """Build from mlp.""" if ep_size > 1: + assert fp32_acc is False, 'FP32 MoE reduction is not supported by the DeepEP backend yet.' + assert output_scale == 1.0, 'MoE output scaling is not supported by the DeepEP backend yet.' assert custom_gateup_act is False, 'Custom gate up activation is not supported in EP MoE.' return FusedDeepEpMoEBlockedF8Impl(ep_size=ep_size, ep_group=ep_group, @@ -479,7 +532,28 @@ def build(top_k: int, fp8_dtype=fp8_dtype, num_max_dispatch_tokens_per_rank=num_max_dispatch_tokens_per_rank, layer_idx=layer_idx) + elif use_deep_gemm: + try: + import deep_gemm # noqa: F401 + except ImportError as e: + raise ImportError('The DeepGEMM MoE path requires the installable deep_gemm package.') from e + return FusedMoENormal(ep_size=1, + ep_group=ep_group, + num_experts=num_experts, + hidden_dim=hidden_dim, + renormalize=renormalize, + block_size=block_size, + top_k=top_k, + out_dtype=out_dtype, + fp8_dtype=fp8_dtype, + scale_fmt=None, + fp32_acc=fp32_acc, + output_scale=output_scale, + num_max_dispatch_tokens_per_rank=num_max_dispatch_tokens_per_rank, + layer_index=layer_idx) else: + if fp32_acc or output_scale != 1.0: + raise ValueError('FP32 MoE reduction and output scaling require use_deep_gemm=True.') return TritonFusedMoEBlockedF8Impl(top_k=top_k, num_experts=num_experts, renormalize=renormalize, diff --git a/lmdeploy/pytorch/backends/cuda/op_backend.py b/lmdeploy/pytorch/backends/cuda/op_backend.py index 53e71239a1..d87ba78999 100644 --- a/lmdeploy/pytorch/backends/cuda/op_backend.py +++ b/lmdeploy/pytorch/backends/cuda/op_backend.py @@ -111,6 +111,9 @@ def get_layer_impl_builder(cls, layer_type: OpType): elif layer_type == OpType.GatedDeltaRule: from .gated_delta_rule import CudaGatedDeltaRuleBuilder return CudaGatedDeltaRuleBuilder + elif layer_type == OpType.Kda: + from .kda import CudaKdaBuilder + return CudaKdaBuilder elif layer_type == OpType.CacheBlockCopy: from .cache_block_copy import CudaCacheBlockCopyBuilder return CudaCacheBlockCopyBuilder diff --git a/lmdeploy/pytorch/backends/kda.py b/lmdeploy/pytorch/backends/kda.py new file mode 100644 index 0000000000..63004e23ca --- /dev/null +++ b/lmdeploy/pytorch/backends/kda.py @@ -0,0 +1,39 @@ +# Copyright (c) OpenMMLab. All rights reserved. +from abc import ABC, abstractmethod +from typing import Any + +import torch + + +class KdaImpl(ABC): + """Backend interface for Kimi Delta Attention inference.""" + + @abstractmethod + def forward( + self, + mixed_qkv: torch.Tensor, + raw_gate: torch.Tensor, + raw_beta: torch.Tensor, + conv_weight: torch.Tensor, + conv_bias: torch.Tensor | None, + a_log: torch.Tensor, + dt_bias: torch.Tensor, + conv_state: torch.Tensor, + recurrent_state: torch.Tensor, + metadata: Any, + num_heads: int, + head_dim: int, + lower_bound: float, + ) -> torch.Tensor: + """Run short convolution followed by the KDA recurrence.""" + raise NotImplementedError + + +class KdaBuilder: + """Build a device-specific KDA implementation.""" + + @staticmethod + @abstractmethod + def build() -> KdaImpl: + """Build implementation.""" + raise NotImplementedError diff --git a/lmdeploy/pytorch/backends/moe.py b/lmdeploy/pytorch/backends/moe.py index a6b72d3c3b..1871ed4002 100644 --- a/lmdeploy/pytorch/backends/moe.py +++ b/lmdeploy/pytorch/backends/moe.py @@ -264,6 +264,9 @@ def build(top_k: int, fp8_dtype: torch.dtype = torch.float8_e4m3fn, num_max_dispatch_tokens_per_rank: int = 128, layer_idx: int = 0, - custom_gateup_act: bool = False): + custom_gateup_act: bool = False, + fp32_acc: bool = False, + output_scale: float = 1.0, + use_deep_gemm: bool = False): """Build from mlp.""" raise NotImplementedError diff --git a/lmdeploy/pytorch/config.py b/lmdeploy/pytorch/config.py index 8b0e93510d..7489e766ca 100644 --- a/lmdeploy/pytorch/config.py +++ b/lmdeploy/pytorch/config.py @@ -474,6 +474,12 @@ class ModelConfig: use_standard_kv_cache: bool = True post_build_func: Callable[['ModelConfig', int], None] | None = None + # A model can reject generic prefix-cache reuse when its auxiliary state + # cannot be restored at scheduler block boundaries. ``None`` keeps the + # default supported behavior; a non-empty reason is surfaced before model + # weights are built. + prefix_caching_unsupported_reason: str | None = None + # check env for model-device combination check_env_func: Callable = _default_check_env @@ -493,6 +499,10 @@ class ModelConfig: # Number of contiguous TP ranks that own the same logical KV-head shard. num_replicate_key_value_heads: int = 1 + # Model-specific defaults that must be present before the distributed + # process group is initialized. Explicit process environment values win. + process_group_env_defaults: dict[str, str] = field(default_factory=dict) + @property def use_mla_fp8_cache(self): """Whether MLA uses the DeepSeek-V3.2 FP8 cache layout.""" diff --git a/lmdeploy/pytorch/configurations/glm5_next.py b/lmdeploy/pytorch/configurations/glm5_next.py new file mode 100644 index 0000000000..c465aabc35 --- /dev/null +++ b/lmdeploy/pytorch/configurations/glm5_next.py @@ -0,0 +1,311 @@ +# Copyright (c) OpenMMLab. All rights reserved. +"""PyTorch engine configuration for GLM-5.3-Flash.""" + +import torch + +from lmdeploy.pytorch.config import BlockCacheSpec, ModelConfig, StateCacheSpec +from lmdeploy.pytorch.consts import ( + DSA_INDEXER_K_CACHE_NAME, + GLM5_KDA_CONV_STATE, + GLM5_KDA_RECURRENT_STATE, + GLM5_KPOOL_TAIL_K_STATE, + GLM5_KPOOL_TAIL_SCORE_STATE, + dsa_packed_indexer_k_cache_shape, +) + +from .builder import AutoModelConfigBuilder +from .deepseek_v32 import DeepseekV32ModelConfigBuilder +from .qwen3_next import _check_env_qwen3_next + +_GLM5_LINEAR_LAYER_TYPE = 'linear_attention' +_GLM5_FULL_LAYER_TYPES = frozenset({ + 'deepseek_sparse_attention', + 'full_attention', +}) + + +def is_glm5_kda_layer(text_config, layer_idx: int) -> bool: + """Return the normalized hybrid-layer kind across process boundaries. + + ``_resolve_glm5_linear_config`` materializes ``linear_layer_ids`` as an + instance field, so it survives Ray serialization. In contrast, the + compatibility method installed on a native Transformers config class is + process-local and is not available in a freshly spawned worker. + """ + linear_layer_ids = getattr(text_config, 'linear_layer_ids', None) + if linear_layer_ids is None: + _, linear_layer_ids, _ = _resolve_glm5_linear_config(text_config) + return int(layer_idx) in linear_layer_ids + + +def _compat_is_kda_layer(text_config, layer_idx: int) -> bool: + """Provide the legacy predicate for native Transformers configs.""" + return is_glm5_kda_layer(text_config, layer_idx) + + +def _set_layer_ids_compat(text_config, name: str, + layer_ids: list[int]) -> None: + """Expose a resolved layer map without breaking legacy properties.""" + try: + setattr(text_config, name, layer_ids) + except AttributeError: + # The LMDeploy-owned legacy config exposes these as read-only + # properties. Its properties read the normalized legacy dict below. + existing = list(getattr(text_config, name)) + if existing != layer_ids: + raise ValueError( + f'GLM-5.3 {name}={existing} conflicts with layer_types ' + f'({layer_ids}).') + + +def _resolve_glm5_linear_config(text_config): + """Normalize official and legacy GLM-5.3 hybrid-layer schemas. + + Transformers 5.16 represents the hybrid layout with ``layer_types`` and + linear-attention geometry with scalar fields. Older checkpoints and the + LMDeploy fallback config expose ``linear_attn_config`` plus convenience + layer-id properties. Keep ``layer_types`` authoritative when present and + materialize the legacy view consumed by the existing model components. + """ + num_layers = int(text_config.num_hidden_layers) + legacy = dict(getattr(text_config, 'linear_attn_config', None) or {}) + layer_types = getattr(text_config, 'layer_types', None) + + if layer_types is not None: + layer_types = list(layer_types) + if len(layer_types) != num_layers: + raise ValueError( + 'GLM-5.3 requires one layer_type per hidden layer, but got ' + f'{len(layer_types)} layer types for {num_layers} layers.') + invalid_types = sorted( + set(layer_types).difference({_GLM5_LINEAR_LAYER_TYPE}, + _GLM5_FULL_LAYER_TYPES)) + if invalid_types: + raise ValueError( + 'GLM-5.3 layer_types only supports linear_attention, ' + 'deepseek_sparse_attention, or full_attention, but got ' + f'{invalid_types}.') + linear_layer_ids = [ + layer_idx for layer_idx, layer_type in enumerate(layer_types) + if layer_type == _GLM5_LINEAR_LAYER_TYPE + ] + full_attention_layer_ids = [ + layer_idx for layer_idx, layer_type in enumerate(layer_types) + if layer_type != _GLM5_LINEAR_LAYER_TYPE + ] + else: + linear_layer_ids = getattr(text_config, 'linear_layer_ids', None) + full_attention_layer_ids = getattr( + text_config, 'full_attention_layer_ids', None) + if linear_layer_ids is None: + linear_layer_ids = legacy.get('kda_layers') + if full_attention_layer_ids is None: + full_attention_layer_ids = legacy.get('full_attn_layers') + + all_layer_ids = set(range(num_layers)) + if linear_layer_ids is None and full_attention_layer_ids is None: + # Match the original GLM-5.3 fallback config for checkpoints that + # predate both explicit schemas. + linear_layer_ids = [ + layer_idx for layer_idx in range(num_layers) + if layer_idx % 4 != 3 + ] + if linear_layer_ids is None: + linear_layer_ids = sorted( + all_layer_ids.difference(full_attention_layer_ids)) + if full_attention_layer_ids is None: + full_attention_layer_ids = sorted( + all_layer_ids.difference(linear_layer_ids)) + linear_layer_ids = list(linear_layer_ids) + full_attention_layer_ids = list(full_attention_layer_ids) + + expected_layer_ids = list(range(num_layers)) + resolved_layer_ids = linear_layer_ids + full_attention_layer_ids + if sorted(resolved_layer_ids) != expected_layer_ids: + raise ValueError( + 'GLM-5.3 linear/full layer ids must form an exact partition of ' + f'[0, {num_layers}), got linear={linear_layer_ids}, ' + f'full={full_attention_layer_ids}.') + + def _linear_value(native_name: str, legacy_name: str, default=None): + value = getattr(text_config, native_name, None) + if value is None: + value = legacy.get(legacy_name, default) + return value + + num_heads = _linear_value('linear_num_heads', 'num_heads') + head_dim = _linear_value('linear_head_dim', 'head_dim') + conv_kernel_size = _linear_value('linear_conv_kernel_dim', + 'short_conv_kernel_size') + gate_lower_bound = _linear_value('linear_lower_bound', + 'gate_lower_bound', -5.0) + missing = [ + name for name, value in ( + ('linear_num_heads/num_heads', num_heads), + ('linear_head_dim/head_dim', head_dim), + ('linear_conv_kernel_dim/short_conv_kernel_size', + conv_kernel_size), + ) if value is None + ] + if missing: + raise ValueError('GLM-5.3 is missing linear-attention config fields: ' + + ', '.join(missing) + '.') + + legacy.update({ + 'num_heads': int(num_heads), + 'head_dim': int(head_dim), + 'short_conv_kernel_size': int(conv_kernel_size), + 'gate_lower_bound': gate_lower_bound, + 'kda_layers': linear_layer_ids, + 'full_attn_layers': full_attention_layer_ids, + }) + text_config.linear_attn_config = legacy + _set_layer_ids_compat(text_config, 'linear_layer_ids', linear_layer_ids) + _set_layer_ids_compat(text_config, 'full_attention_layer_ids', + full_attention_layer_ids) + if not callable(getattr(text_config, 'is_kda_layer', None)): + # Install a class method instead of an instance closure so HF's + # ``to_dict`` remains JSON serializable after model-config building. + setattr(type(text_config), 'is_kda_layer', _compat_is_kda_layer) + return legacy, linear_layer_ids, full_attention_layer_ids + + +def _check_env_glm5_next(device: str): + """Validate the reused KDA, dense-MLA, and FA3 dependencies.""" + _check_env_qwen3_next(device) + if device != 'cuda': + return + + try: + import flash_mla + except ImportError: + raise ImportError('GLM-5.3 CUDA support requires .') + if not hasattr(flash_mla, 'flash_mla_with_kvcache'): + raise RuntimeError( + 'GLM-5.3 requires FlashMLA dense paged-attention support.') + + # KPool group selection reuses LMDeploy's shared sparse-index Top-K + # implementation; validate the two model geometries without importing + # SGLang providers. + from lmdeploy.pytorch.kernels.cuda.sparse_index_topk import ( + is_sparse_index_topk_supported, + ) + if not all(is_sparse_index_topk_supported(k) for k in (512, 2048)): + raise RuntimeError( + 'GLM-5.3 KPool requires LMDeploy sparse-index Top-K support for ' + 'K=512 and K=2048.') + + +def _finalize_glm5_cache_specs(model_config: ModelConfig, + block_size: int) -> None: + """Materialize the independent pooled-index cache at page size 64.""" + text_config = model_config.llm_config + if block_size != 64: + raise ValueError( + 'GLM-5.3 KPool requires block_size=64, ' + f'got block_size={block_size}.') + model_config.cache_shapes = [] + model_config.block_cache_specs = [ + BlockCacheSpec( + DSA_INDEXER_K_CACHE_NAME, + # Standard MLA caches are already compacted to the 11 full + # attention layers. Named cache rows use that same local id + # space; global decoder ids are translated by the model. + list(range(model_config.num_layers)), + dsa_packed_indexer_k_cache_shape( + block_size, text_config.index_head_dim), + torch.uint8, + ) + ] + + +class Glm5NextModelConfigBuilder(AutoModelConfigBuilder): + """Combine sparse-MLA KV cache and KDA recurrent-state resources.""" + + @classmethod + def condition(cls, hf_config): + return getattr(hf_config, 'model_type', None) == 'glm5_next' + + @classmethod + def build(cls, hf_config, model_path: str | None = None, **kwargs): + if not hasattr(hf_config, 'text_config'): + raise ValueError('GLM-5.3 config must define `text_config`.') + + text_config = hf_config.text_config + quant_config = getattr(hf_config, 'quantization_config', None) + if quant_config is not None: + text_config.quantization_config = quant_config + + linear_config, linear_layer_ids, full_attention_layer_ids = ( + _resolve_glm5_linear_config(text_config)) + config = DeepseekV32ModelConfigBuilder.build( + text_config, model_path=model_path, **kwargs) + + tp = kwargs.get('tp', 1) + device_type = kwargs.get('device_type', 'auto') + if device_type == 'cuda' and tp > 1: + config.process_group_env_defaults.setdefault( + 'NCCL_NVLS_ENABLE', '0') + num_linear_layers = len(linear_layer_ids) + num_full_layers = len(full_attention_layer_ids) + num_heads = linear_config['num_heads'] + head_dim = linear_config['head_dim'] + conv_kernel_size = linear_config['short_conv_kernel_size'] + if num_heads % tp: + raise ValueError( + f'GLM-5.3 linear attention has {num_heads} heads, which is ' + f'not divisible by TP={tp}.') + + local_heads = num_heads // tp + conv_dim = 3 * local_heads * head_dim + config.num_layers = num_full_layers + # GLM owns KPool selection in the model adapter, rather than the + # DeepSeek-V3.2 token indexer selected by mla_index_topk. Keeping this + # unset also preserves the BF16 latent MLA cache policy. + config.mla_index_topk = None + config.k_head_dim = text_config.kv_lora_rank + 64 + config.cache_shapes = [] + config.block_cache_specs = [] + config.post_build_func = _finalize_glm5_cache_specs + config.state_cache_specs = [ + StateCacheSpec( + GLM5_KDA_CONV_STATE, + (num_linear_layers, conv_dim, conv_kernel_size), + torch.bfloat16, + ), + StateCacheSpec( + GLM5_KDA_RECURRENT_STATE, + (num_linear_layers, local_heads, head_dim, head_dim), + torch.float32, + ), + StateCacheSpec( + GLM5_KPOOL_TAIL_K_STATE, + (num_full_layers, text_config.index_kpool, + text_config.index_head_dim), + torch.bfloat16, + ), + StateCacheSpec( + GLM5_KPOOL_TAIL_SCORE_STATE, + (num_full_layers, text_config.index_kpool, + text_config.index_head_dim), + torch.bfloat16, + ), + ] + # Scheduler compatibility: named specs own allocation, while this + # bridge keeps hybrid state checkpointing enabled in legacy callers. + config.states_shapes = [ + (tuple(spec.shape), spec.dtype) + for spec in config.state_cache_specs + ] + config.is_gated_delta = True + config.prefix_caching_unsupported_reason = ( + 'GLM-5.3 KPool packs 4-token groups into 256-token owner blocks, ' + 'so scheduler prefix blocks cannot restore exact KPool state.') + config.check_env_func = _check_env_glm5_next + config.hf_config = hf_config + config.llm_config = text_config + + text_dtype = getattr(text_config, 'dtype', None) + if text_dtype is not None: + hf_config.dtype = text_dtype + return config diff --git a/lmdeploy/pytorch/consts.py b/lmdeploy/pytorch/consts.py index 92a37ac45a..82fe4cff79 100644 --- a/lmdeploy/pytorch/consts.py +++ b/lmdeploy/pytorch/consts.py @@ -14,6 +14,14 @@ DSA_INDEXER_K_CACHE_NAME = 'dsa_indexer_k' DSA_INDEX_SCALE_BYTES = 4 +# GLM-5.3 hybrid sequence-state resources. These names are shared by its +# config declaration and model consumer so they cannot silently drift back to +# anonymous positional state caches. +GLM5_KDA_CONV_STATE = 'glm5_kda_conv' +GLM5_KDA_RECURRENT_STATE = 'glm5_kda_recurrent' +GLM5_KPOOL_TAIL_K_STATE = 'glm5_kpool_tail_k' +GLM5_KPOOL_TAIL_SCORE_STATE = 'glm5_kpool_tail_score' + def v4_packed_index_cache_shape(entries_per_block: int, head_dim: int) -> tuple[int, int, int]: """Return the logical uint8 shape for the packed V4 index cache.""" diff --git a/lmdeploy/pytorch/engine/executor/base.py b/lmdeploy/pytorch/engine/executor/base.py index 5766a4c091..bda66b9901 100644 --- a/lmdeploy/pytorch/engine/executor/base.py +++ b/lmdeploy/pytorch/engine/executor/base.py @@ -63,6 +63,12 @@ def _maybe_disable_unsupported_prefix_caching(self, *, check_window: bool = True """Disable prefix caching for unsupported executor/cache modes.""" if not getattr(self.cache_config, 'enable_prefix_caching', False): return + unsupported_reason = getattr( + self.model_config, 'prefix_caching_unsupported_reason', None) + if unsupported_reason: + raise ValueError( + 'Prefix caching is not supported for this model: ' + f'{unsupported_reason} Set enable_prefix_caching=False.') if check_window and self.cache_config.window_size is not None and self.cache_config.window_size > 0: # do not support generic sliding window prefix caching logger.warning('Sliding window prefix caching is not supported.') diff --git a/lmdeploy/pytorch/engine/executor/base_worker.py b/lmdeploy/pytorch/engine/executor/base_worker.py index a33ae5fb6f..70daf1f34d 100644 --- a/lmdeploy/pytorch/engine/executor/base_worker.py +++ b/lmdeploy/pytorch/engine/executor/base_worker.py @@ -1,6 +1,7 @@ # Copyright (c) OpenMMLab. All rights reserved. import asyncio import gc +import os from typing import Any from lmdeploy.pytorch.backends.selector import get_backend @@ -58,6 +59,9 @@ def __init__( def init_process_group(self, rank: int, master_addr: str = None, master_port: str = None): """Initialize process group.""" + for key, value in self.model_config.process_group_env_defaults.items(): + os.environ.setdefault(key, value) + self.rank = rank if self.world_size > 1: if master_addr is not None and master_port is not None: diff --git a/lmdeploy/pytorch/engine/executor/ray_executor.py b/lmdeploy/pytorch/engine/executor/ray_executor.py index 5f749a7fa6..a84b9cd65b 100644 --- a/lmdeploy/pytorch/engine/executor/ray_executor.py +++ b/lmdeploy/pytorch/engine/executor/ray_executor.py @@ -73,11 +73,16 @@ def _update_env_cuda_alloc_conf(env_vars: dict): env_vars['PYTORCH_CUDA_ALLOC_CONF'] = cuda_alloc_conf -def _update_runtime_envs(runtime_env: dict): +def _update_runtime_envs( + runtime_env: dict, + process_group_env_defaults: dict[str, str] | None = None, +): """Update runtime envs.""" new_envs = _envs.get_all_envs() env_vars: dict = runtime_env.get('env_vars', {}) env_vars.update(new_envs) + for key, default in (process_group_env_defaults or {}).items(): + env_vars.setdefault(key, os.environ.get(key, default)) _update_env_cuda_alloc_conf(env_vars) runtime_env['env_vars'] = env_vars return runtime_env @@ -661,6 +666,8 @@ def get_priority(ip): def _init_workers_ray(self, placement_group: PlacementGroup, worker_kwargs: dict): """Init worker ray.""" device_str = get_device_str() + process_group_env_defaults = ( + worker_kwargs['model_config'].process_group_env_defaults) bundle_indices = [] if not _envs.ray_external_pg_bundles: for bundle_id, bundle in enumerate(placement_group.bundle_specs): @@ -697,7 +704,8 @@ def _init_workers_ray(self, placement_group: PlacementGroup, worker_kwargs: dict if device_str == 'GPU': runtime_env = dict() - runtime_env = _update_runtime_envs(runtime_env) + runtime_env = _update_runtime_envs( + runtime_env, process_group_env_defaults) if self._try_symm_mem: # Symmetric-memory IPC needs peer TP GPUs to stay visible. # Keep the inherited visibility and bind each actor below. @@ -712,7 +720,8 @@ def _init_workers_ray(self, placement_group: PlacementGroup, worker_kwargs: dict )(RayWorkerWrapper).remote(**worker_kwargs) else: runtime_env = dict() - runtime_env = _update_runtime_envs(runtime_env) + runtime_env = _update_runtime_envs( + runtime_env, process_group_env_defaults) worker = ray.remote( num_cpus=0, num_gpus=0, diff --git a/lmdeploy/pytorch/kernels/cuda/activation.py b/lmdeploy/pytorch/kernels/cuda/activation.py index 6eebe91af7..c198e4999e 100644 --- a/lmdeploy/pytorch/kernels/cuda/activation.py +++ b/lmdeploy/pytorch/kernels/cuda/activation.py @@ -2,6 +2,7 @@ import torch import triton import triton.language as tl +from triton.language.extra import libdevice from .utils import get_device_props @@ -19,6 +20,8 @@ def _silu_and_mul_kernel( stride_om: tl.constexpr, stride_on: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, + SWIGLU_LIMIT: tl.constexpr, + PRECISE_MUL: tl.constexpr, ): """Silu and mul kernel.""" n_block_id = tl.program_id(0) @@ -40,11 +43,19 @@ def _silu_and_mul_kernel( for _ in tl.range(m_id_start, M, m_id_stride): gate = tl.load(gate_ptrs, mask=mask) up = tl.load(up_ptrs, mask=mask) + if SWIGLU_LIMIT is not None: + gate = tl.minimum(gate, SWIGLU_LIMIT) + up = tl.maximum(tl.minimum(up, SWIGLU_LIMIT), -SWIGLU_LIMIT) # exp expect fp32 gate = gate.to(tl.float32) - gate = gate / (1 + fast_expf(-gate)) - gate = gate.to(gateup_ptr.dtype.element_ty) + if PRECISE_MUL: + exp_neg_gate = libdevice.exp(-gate) + else: + exp_neg_gate = fast_expf(-gate) + gate = gate / (1 + exp_neg_gate) + if not PRECISE_MUL: + gate = gate.to(gateup_ptr.dtype.element_ty) out = gate * up tl.store(out_ptrs, out, mask=mask) @@ -54,7 +65,10 @@ def _silu_and_mul_kernel( out_ptrs += m_id_stride * stride_om -def silu_and_mul(gate_up: torch.Tensor, out: torch.Tensor = None): +def silu_and_mul(gate_up: torch.Tensor, + out: torch.Tensor = None, + swiglu_limit: float | None = None, + precise_mul: bool = False): """Silu and mul.""" assert gate_up.dim() == 2 @@ -85,6 +99,8 @@ def silu_and_mul(gate_up: torch.Tensor, out: torch.Tensor = None): stride_om=out.stride(0), stride_on=out.stride(1), BLOCK_SIZE_N=BLOCK_SIZE_N, + SWIGLU_LIMIT=swiglu_limit, + PRECISE_MUL=precise_mul, num_warps=num_warps, num_stages=num_stages) diff --git a/lmdeploy/pytorch/kernels/cuda/moe/ep.py b/lmdeploy/pytorch/kernels/cuda/moe/ep.py index ed51f840a8..f5d9347c12 100644 --- a/lmdeploy/pytorch/kernels/cuda/moe/ep.py +++ b/lmdeploy/pytorch/kernels/cuda/moe/ep.py @@ -142,13 +142,13 @@ def _fwd_kernel_ep_gather( output_tensor_stride1, topk_num: tl.constexpr, BLOCK_D: tl.constexpr, + FP32_ACC: tl.constexpr, + OUTPUT_SCALE: tl.constexpr, ): cur_block = tl.program_id(0) start_cur_token = tl.program_id(1) grid_num = tl.num_programs(1) - # align with xtuner rl - compute_dtype = output_tensor.dtype.element_ty - # compute_dtype = tl.float32 + compute_dtype = tl.float32 if FP32_ACC else output_tensor.dtype.element_ty for cur_token in range(start_cur_token, total_token_num, grid_num): off_d = tl.arange(0, BLOCK_D) @@ -160,6 +160,7 @@ def _fwd_kernel_ep_gather( acc_weight = tl.load(recv_topk_weight + cur_token * recv_topk_weight_stride0 + topk_index) tmp = tl.load(input_tensor + source_token_index * input_tensor_stride0 + cur_block * BLOCK_D + off_d) accumulator += tmp.to(compute_dtype) * acc_weight.to(compute_dtype) + accumulator *= OUTPUT_SCALE tl.store( output_tensor + cur_token * output_tensor_stride0 + cur_block * BLOCK_D + off_d, accumulator.to(output_tensor.dtype.element_ty), @@ -173,6 +174,8 @@ def ep_gather( recv_topk_weight: torch.Tensor, input_index: torch.Tensor, output_tensor: torch.Tensor, + fp32_acc: bool = False, + output_scale: float = 1.0, ): BLOCK_D = 1024 # block size of quantization num_warps = 2 @@ -200,6 +203,8 @@ def ep_gather( topk_num=recv_topk_ids.shape[1], num_warps=num_warps, BLOCK_D=BLOCK_D, + FP32_ACC=fp32_acc, + OUTPUT_SCALE=output_scale, ) return diff --git a/lmdeploy/pytorch/kernels/cuda/moe/ep_fp8.py b/lmdeploy/pytorch/kernels/cuda/moe/ep_fp8.py index 8eab01008f..374d69f252 100644 --- a/lmdeploy/pytorch/kernels/cuda/moe/ep_fp8.py +++ b/lmdeploy/pytorch/kernels/cuda/moe/ep_fp8.py @@ -145,50 +145,81 @@ def _deepgemm_grouped_fp8_nt_contiguous(input_tuple, w_tuple, out: torch.Tensor, return deep_gemm.m_grouped_fp8_gemm_nt_contiguous(input_tuple, w_tuple, out, m_indices) +def _get_compact_all_tokens(num_assignments: int, num_experts: int, block_e: int = 128) -> int: + """Maximum expert-aligned rows for a graph-stable compact layout.""" + max_nonempty_experts = min(num_assignments, num_experts) + return block_e * (max_nonempty_experts + (num_assignments - max_nonempty_experts) // block_e) + + def fused_moe_v3_fp8( hidden_states_fp8: tuple[torch.Tensor, torch.Tensor], topk_idx, topk_weights, w13_weight_fp8: tuple[torch.Tensor, torch.Tensor], w2_weight_fp8: tuple[torch.Tensor, torch.Tensor], - num_recv_tokens_per_expert: list[int] | None, + num_recv_tokens_per_expert: list[int] | torch.Tensor | None, + *, + compact_layout: bool = False, + act_func=None, + fp32_acc: bool = False, + output_scale: float = 1.0, ): hidden_states_fp8, hidden_states_scale = hidden_states_fp8 if num_recv_tokens_per_expert is None: return hidden_states_fp8.to(torch.bfloat16) - all_tokens = sum(num_recv_tokens_per_expert) + if compact_layout: + assert isinstance(num_recv_tokens_per_expert, torch.Tensor) + all_tokens = _get_compact_all_tokens(topk_idx.numel(), num_recv_tokens_per_expert.numel()) + num_recv_tokens_per_expert_gpu = num_recv_tokens_per_expert.to(device=hidden_states_fp8.device, + dtype=torch.int32) + num_recv_tokens_per_expert_gpu = (num_recv_tokens_per_expert_gpu + 127) // 128 * 128 + num_recv_tokens_per_expert_gpu[-1].add_(all_tokens - num_recv_tokens_per_expert_gpu.sum()) + else: + all_tokens = sum(num_recv_tokens_per_expert) + num_recv_tokens_per_expert_gpu = torch.tensor(num_recv_tokens_per_expert, + dtype=torch.int32, + pin_memory=True, + device='cpu').cuda(non_blocking=True) if all_tokens <= 0: return hidden_states_fp8.to(torch.bfloat16) - from lmdeploy.pytorch.third_party.deep_gemm import get_mn_major_tma_aligned_tensor m, k = hidden_states_fp8.size() n = w13_weight_fp8[0].size(1) block_size = k // hidden_states_scale.size(1) gather_out = torch.empty_like(hidden_states_fp8, device=hidden_states_fp8.device, dtype=torch.bfloat16) - input_tensor = torch.empty((all_tokens, k), device=hidden_states_fp8.device, dtype=hidden_states_fp8.dtype) - input_tensor_scale = torch.empty((all_tokens, k // block_size), - device=hidden_states_fp8.device, - dtype=torch.float32) + # Padding keeps its expert id in the existing scatter kernel. Zero its + # inputs/scales; gather reads only the real routed rows via output_index. + allocator = torch.zeros if compact_layout else torch.empty + input_tensor = allocator((all_tokens, k), device=hidden_states_fp8.device, dtype=hidden_states_fp8.dtype) + input_tensor_scale = allocator((all_tokens, k // block_size), + device=hidden_states_fp8.device, + dtype=torch.float32) m_indices = torch.empty(all_tokens, device=hidden_states_fp8.device, dtype=torch.int32) output_index = torch.empty_like(topk_idx) - num_recv_tokens_per_expert_gpu = torch.tensor(num_recv_tokens_per_expert, - dtype=torch.int32, - pin_memory=True, - device='cpu').cuda(non_blocking=True) expert_start_loc = torch.empty_like(num_recv_tokens_per_expert_gpu) ep_scatter_fp8(hidden_states_fp8, hidden_states_scale, topk_idx, num_recv_tokens_per_expert_gpu, expert_start_loc, input_tensor, input_tensor_scale, m_indices, output_index) del hidden_states_fp8 + from lmdeploy.pytorch.third_party.deep_gemm import get_mn_major_tma_aligned_tensor gateup_output = torch.empty((all_tokens, n), device=gather_out.device, dtype=torch.bfloat16) input_tensor_scale = get_mn_major_tma_aligned_tensor(input_tensor_scale) _deepgemm_grouped_fp8_nt_contiguous((input_tensor, input_tensor_scale), w13_weight_fp8, gateup_output, m_indices) - down_input = torch.empty((all_tokens, n // 2), device=gateup_output.device, dtype=torch.bfloat16) - silu_and_mul(gateup_output.view(-1, n), down_input) + if act_func is None: + down_input = torch.empty((all_tokens, n // 2), device=gateup_output.device, dtype=torch.bfloat16) + silu_and_mul(gateup_output.view(-1, n), down_input) + else: + down_input = act_func(gateup_output.view(-1, n)) del gateup_output down_input_fp8, down_input_scale = per_token_group_quant_fp8(down_input, block_size) down_input_scale = get_mn_major_tma_aligned_tensor(down_input_scale) down_output = torch.empty((all_tokens, k), device=gather_out.device, dtype=torch.bfloat16) _deepgemm_grouped_fp8_nt_contiguous((down_input_fp8, down_input_scale), w2_weight_fp8, down_output, m_indices) - ep_gather(down_output, topk_idx, topk_weights, output_index, gather_out) + ep_gather(down_output, + topk_idx, + topk_weights, + output_index, + gather_out, + fp32_acc=fp32_acc, + output_scale=output_scale) return gather_out diff --git a/lmdeploy/pytorch/kernels/cuda/moe/route_noaux_tc.py b/lmdeploy/pytorch/kernels/cuda/moe/route_noaux_tc.py index acb2ca5d60..b7ca8e3440 100644 --- a/lmdeploy/pytorch/kernels/cuda/moe/route_noaux_tc.py +++ b/lmdeploy/pytorch/kernels/cuda/moe/route_noaux_tc.py @@ -99,7 +99,10 @@ def fused_noaux_tc_routing( bias = bias.float().contiguous() topk_weight = torch.empty(batch_size, top_k, device=logits.device, dtype=torch.float32) topk_idx = torch.empty(batch_size, top_k, device=logits.device, dtype=torch.int64) - block_size = num_experts + # Triton requires every ``tl.arange`` extent to be a power of two. The + # ungrouped GLM case has 288 experts, so pad its masked expert axis. + block_size = (triton.next_power_of_2(num_experts) + if n_group == 1 else num_experts) assert block_size % 32 == 0, 'num_experts must be a multiple of 32 for optimal performance' grid = (batch_size, ) _fused_noaux_tc_kernel[grid]( diff --git a/lmdeploy/pytorch/kernels/cuda/sparse_mla_tilelang.py b/lmdeploy/pytorch/kernels/cuda/sparse_mla_tilelang.py new file mode 100644 index 0000000000..3328b70b49 --- /dev/null +++ b/lmdeploy/pytorch/kernels/cuda/sparse_mla_tilelang.py @@ -0,0 +1,234 @@ +# Copyright (c) OpenMMLab. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""TileLang sparse MLA decode attention for CUDA BF16 tensors. + +The kernel schedule in this file is adapted from SGLang's +sparse_attention_fwd_kernel_v1 at commit 9e692c9216c3, distributed under the +Apache License 2.0: + +https://github.com/sgl-project/sglang/blob/9e692c9216c3/python/sglang/kernels/ops/attention/dsa/tilelang_kernel.py + +Only the CUDA BF16 path used by GLM-5.3-Flash is retained. The public wrapper +uses generic sparse-MLA names and has no SGLang runtime dependency. +""" + +import tilelang +import tilelang.language as T +import torch + +tilelang.set_log_level('WARNING') + + +@tilelang.jit( + out_idx=[-1], + pass_configs={ + tilelang.PassConfigKey.TL_DISABLE_TMA_LOWER: True, + tilelang.PassConfigKey.TL_DISABLE_WARP_SPECIALIZED: True, + }, +) +def _sparse_mla_bf16_fwd_kernel( + num_heads, + dim, + topk, + *, + storage_dim=None, + kv_group=1, + sm_scale=None, + is_causal=True, + block_I=64, + num_stages=2, + threads=256, +): + if storage_dim is None: + storage_dim = dim + assert storage_dim >= dim + assert ( + dim == tilelang.math.next_power_of_2(dim) or dim % 64 == 0 + ), f"dim={dim} must be a power of 2 or a multiple of 64" + assert is_causal, "non-causal is not supported" + assert ( + topk % block_I == 0 + ), "otherwise will load some index=0 thus causing wrong kv to be loaded" + if sm_scale is None: + sm_scale = (1.0 / dim) ** 0.5 * 1.44269504 # log2(e) + else: + sm_scale = sm_scale * 1.44269504 # log2(e) + + batch = T.symbolic("batch") + seq_len = T.symbolic("seq_len") + seq_len_kv = T.symbolic("seq_len_kv") + + head_kv = num_heads // kv_group + q_shape = [batch, seq_len, num_heads, dim] + kv_shape = [batch, seq_len_kv, kv_group, storage_dim] + o_shape = [batch, seq_len, num_heads, dim] + indices_shape = [batch, seq_len, kv_group, topk] + indices_dtype = "int32" + dtype = "bfloat16" + accum_dtype = "float" + + H = head_kv + padded_H = max(tilelang.math.next_power_of_2(head_kv), 16) + if padded_H != H: + assert kv_group == 1 + BI = block_I + NI = tilelang.cdiv(topk, block_I) + D = dim + + if head_kv > 64: + assert head_kv % 64 == 0, "head_kv should be a multiple of 64" + REPLICATE_H = head_kv // 64 + else: + REPLICATE_H = 1 + + H_per_block = padded_H if REPLICATE_H == 1 else 64 + + @T.prim_func + def main( + Q: T.Tensor(q_shape, dtype), # type: ignore + KV: T.Tensor(kv_shape, dtype), # type: ignore + Indices: T.Tensor(indices_shape, indices_dtype), # type: ignore + Output: T.Tensor(o_shape, dtype), # type: ignore + ): + with T.Kernel(seq_len * REPLICATE_H, batch, kv_group, threads=threads) as ( + bx, + by, + bz, + ): + Q_shared = T.alloc_shared([H_per_block, D], dtype) + KV_shared = T.alloc_shared([BI, D], dtype) + O_shared = T.alloc_shared([H_per_block, D], dtype) + mask = T.alloc_fragment([BI], "bool") + + acc_o = T.alloc_fragment([H_per_block, D], accum_dtype) + acc_s = T.alloc_fragment([H_per_block, BI], accum_dtype) + S_shared = T.alloc_shared([H_per_block, BI], dtype) + sumexp = T.alloc_fragment([H_per_block], accum_dtype) + sumexp_i = T.alloc_fragment([H_per_block], accum_dtype) + alpha = T.alloc_fragment([H_per_block], accum_dtype) + m_i = T.alloc_fragment([H_per_block], accum_dtype) + m_i_prev = T.alloc_fragment([H_per_block], accum_dtype) + + T.fill(acc_o, 0) + T.fill(sumexp, 0) + T.fill(m_i, -(2**30)) # avoid -inf - inf to cause nan + + b_i, g_i = by, bz + s_i = bx if REPLICATE_H == 1 else (bx // REPLICATE_H) + + H0 = g_i * padded_H + (0 if REPLICATE_H == 1 else (bx % REPLICATE_H) * 64) + H1 = H0 + H_per_block + + T.copy(Q[b_i, s_i, H0:H1, :D], Q_shared) + + for i_i in T.Pipelined(NI, num_stages=num_stages): + + for bi_i in T.Parallel(BI): + mask[bi_i] = Indices[b_i, s_i, g_i, i_i * BI + bi_i] >= 0 + + for bi_i, d_i in T.Parallel(BI, D): + KV_shared[bi_i, d_i] = KV[ + b_i, Indices[b_i, s_i, g_i, i_i * BI + bi_i], g_i, d_i + ] + + for h_i, bi_i in T.Parallel(H_per_block, BI): + acc_s[h_i, bi_i] = T.if_then_else( + mask[bi_i], 0, -T.infinity(acc_s.dtype) + ) + T.gemm( + Q_shared, + KV_shared, + acc_s, + transpose_B=True, + policy=T.GemmWarpPolicy.FullCol, + ) + T.copy(m_i, m_i_prev) + T.reduce_max(acc_s, m_i, dim=1, clear=False) + for h_i in T.Parallel(H_per_block): + alpha[h_i] = T.exp2((m_i_prev[h_i] - m_i[h_i]) * sm_scale) + for h_i, bi_i in T.Parallel(H_per_block, BI): + acc_s[h_i, bi_i] = T.exp2( + acc_s[h_i, bi_i] * sm_scale - m_i[h_i] * sm_scale + ) + T.reduce_sum(acc_s, sumexp_i, dim=1) # is this a accumulate operator? + for h_i in T.Parallel(H_per_block): + sumexp[h_i] = sumexp[h_i] * alpha[h_i] + sumexp_i[h_i] + for h_i, d_i in T.Parallel(H_per_block, D): + acc_o[h_i, d_i] = acc_o[h_i, d_i] * alpha[h_i] + + T.copy(acc_s, S_shared) + T.gemm(S_shared, KV_shared, acc_o, policy=T.GemmWarpPolicy.FullCol) + + # Rescale + for h_i, d_i in T.Parallel(H_per_block, D): + acc_o[h_i, d_i] /= sumexp[h_i] + for h_i in T.Parallel(H_per_block): + sumexp[h_i] = T.log2(sumexp[h_i]) + m_i[h_i] * sm_scale + + T.copy(acc_o, O_shared) + T.copy(acc_o, Output[b_i, s_i, H0:H1, :]) + + return main + + +def sparse_mla_bf16_fwd(q: torch.Tensor, + kv: torch.Tensor, + indices: torch.Tensor, + sm_scale: float) -> torch.Tensor: + """Run sparse MLA decode attention on CUDA BF16 tensors. + + Args: + q: Query tensor with shape [tokens, local_heads, 512]. + kv: Paged latent KV tensor with shape [slots, 1, storage_dim]. The + supported storage widths are 512 and 576; only the first 512 + elements participate in attention. + indices: Physical slot indices with shape [tokens, 1, topk]. + Invalid entries must be -1 and topk must be padded to a + multiple of 64. + sm_scale: Attention scale before the kernel's base-2 conversion. + + Returns: + A BF16 tensor with shape [tokens, local_heads, 512]. + """ + if torch.version.hip is not None: + raise RuntimeError('sparse_mla_bf16_fwd only supports CUDA') + if not (q.is_cuda and kv.is_cuda and indices.is_cuda): + raise ValueError('q, kv, and indices must be CUDA tensors') + if q.dtype != torch.bfloat16 or kv.dtype != torch.bfloat16: + raise TypeError('q and kv must have dtype torch.bfloat16') + if indices.dtype != torch.int32: + raise TypeError('indices must have dtype torch.int32') + if q.ndim != 3 or q.shape[1] <= 0 or q.shape[2] != 512: + raise ValueError( + 'q must have shape [tokens, positive_local_heads, 512], ' + f'got {tuple(q.shape)}') + if kv.ndim != 3 or kv.shape[1] != 1 or kv.shape[2] not in (512, 576): + raise ValueError( + 'kv must have shape [slots, 1, storage_dim] with storage_dim ' + f'in (512, 576), got {tuple(kv.shape)}') + if (indices.ndim != 3 or indices.shape[0] != q.shape[0] + or indices.shape[1] != 1): + raise ValueError( + 'indices must have shape [tokens, 1, topk] matching q, ' + f'got {tuple(indices.shape)}') + if q.device != kv.device or q.device != indices.device: + raise ValueError('q, kv, and indices must be on the same CUDA device') + + topk = indices.shape[-1] + if topk == 0 or topk % 64 != 0: + raise ValueError(f'topk must be a positive multiple of 64, got {topk}') + + if not kv.is_contiguous(): + raise ValueError( + 'kv must be a zero-copy contiguous view of the complete paged cache') + q = q.contiguous() + indices = indices.contiguous() + kernel = _sparse_mla_bf16_fwd_kernel( + num_heads=q.shape[1], + dim=512, + topk=topk, + storage_dim=kv.shape[-1], + sm_scale=sm_scale, + ) + output = kernel(q.unsqueeze(0), kv.unsqueeze(0), indices.unsqueeze(0)) + return output.squeeze(0) diff --git a/lmdeploy/pytorch/models/deepseek_v2.py b/lmdeploy/pytorch/models/deepseek_v2.py index c54f8b52ce..987366312b 100644 --- a/lmdeploy/pytorch/models/deepseek_v2.py +++ b/lmdeploy/pytorch/models/deepseek_v2.py @@ -597,12 +597,15 @@ def __init__(self, config: Any, dtype: torch.dtype = None, device: torch.device = None, - info: EPLBDispatchInfo = None): + info: EPLBDispatchInfo = None, + routed_scaling_factor: float | None = None): super().__init__() self.config = config self.top_k = config.num_experts_per_tok self.n_routed_experts = config.n_routed_experts - self.routed_scaling_factor = config.routed_scaling_factor + if routed_scaling_factor is None: + routed_scaling_factor = config.routed_scaling_factor + self.routed_scaling_factor = routed_scaling_factor self.scoring_func = config.scoring_func self.topk_method = config.topk_method self.n_group = config.n_group @@ -696,6 +699,13 @@ def forward(self, hidden_states: torch.Tensor, routed_experts: torch.Tensor = No class DeepseekV2MoE(nn.Module): """Deepseek v2 MoE.""" + fused_moe_act_func = None + fused_moe_fp32_acc = False + fused_moe_output_scale = 1.0 + fused_moe_use_deep_gemm = False + router_routed_scaling_factor = None + shared_expert_cls = None + def __init__(self, config: Any, layer_idx, @@ -727,9 +737,21 @@ def __init__(self, layer_idx=layer_idx, ) self.num_experts = EPLBManager.num_physical_experts() - self.gate = MoEGate(config, dtype=dtype, device=device, info=eplb_dispatch_info) + self.gate = MoEGate( + config, + dtype=dtype, + device=device, + info=eplb_dispatch_info, + routed_scaling_factor=type(self).router_routed_scaling_factor, + ) else: - self.gate = MoEGate(config, dtype=dtype, device=device, info=None) + self.gate = MoEGate( + config, + dtype=dtype, + device=device, + info=None, + routed_scaling_factor=type(self).router_routed_scaling_factor, + ) self.experts = build_fused_moe( self.hidden_dim, self.ffn_dim, @@ -741,11 +763,16 @@ def __init__(self, all_reduce=moe_all_reduce, quant_config=quantization_config, layer_idx=layer_idx, + act_func=type(self).fused_moe_act_func, + fp32_acc=type(self).fused_moe_fp32_acc, + output_scale=type(self).fused_moe_output_scale, + use_deep_gemm=type(self).fused_moe_use_deep_gemm, ) self.shared_experts = None if config.n_shared_experts is not None: intermediate_size = (config.moe_intermediate_size * config.n_shared_experts) - self.shared_experts = DeepseekV2MLP( + shared_expert_cls = type(self).shared_expert_cls or DeepseekV2MLP + self.shared_experts = shared_expert_cls( config=config, intermediate_size=intermediate_size, dtype=dtype, @@ -1282,9 +1309,9 @@ def __load_kcvc_blocked_fp8(name: str, loaded_weight: torch.Tensor): """Dequant weight.""" if name.endswith('.weight'): weight_name = name - scale_name = name.replace('.weight', '.scale') + scale_name = name.removesuffix('.weight') + '.weight_scale_inv' elif name.endswith('.weight_scale_inv'): - weight_name = name.replace('.weight_scale_inv', '.weight') + weight_name = name.removesuffix('.weight_scale_inv') + '.weight' scale_name = name self._load_buffers[name] = loaded_weight if (weight_name in self._load_buffers and scale_name in self._load_buffers): @@ -1316,7 +1343,11 @@ def __load_kcvc_blocked_fp8(name: str, loaded_weight: torch.Tensor): fp8_quant_scope = quantization_config.get('fp8_quant_scope') loaded_weight = loaded_weight.to(device) - if quant_method == 'fp8' and fp8_quant_scope != 'moe_only': + is_blocked_fp8_tensor = (name.endswith('.weight_scale_inv') + or loaded_weight.dtype == torch.float8_e4m3fn) + if (quant_method == 'fp8' + and fp8_quant_scope != 'moe_only' + and is_blocked_fp8_tensor): # update blocked fp8 weight __load_kcvc_blocked_fp8(name, loaded_weight) else: diff --git a/lmdeploy/pytorch/models/deepseek_v32.py b/lmdeploy/pytorch/models/deepseek_v32.py index 32d1db683b..da5ea71785 100644 --- a/lmdeploy/pytorch/models/deepseek_v32.py +++ b/lmdeploy/pytorch/models/deepseek_v32.py @@ -113,7 +113,8 @@ def _load_fused_qkv_a_weight(name: str, loaded_weight: torch.Tensor, params_dict if param is None: return False - if shard_id == 1 and not name.endswith('.weight_scale_inv'): + if (shard_id == 1 and config.qk_rope_head_dim > 0 + and not name.endswith('.weight_scale_inv')): kv_dim = config.kv_lora_rank + config.qk_rope_head_dim loaded_weight = loaded_weight.to(param.device).unflatten(0, (-1, kv_dim)) rope_weight = loaded_weight[:, config.kv_lora_rank:] @@ -251,6 +252,9 @@ def forward(self, class DeepseekV32Attention(DeepseekV2Attention): + use_sparse_mla = True + mla_head_padding = 0 + def __init__(self, config: Any, layer_idx: int, @@ -344,13 +348,15 @@ def __init__(self, self.softmax_scale = self.softmax_scale * mscale * mscale self.attn_fwd = Attention(self.num_heads, - config.kv_lora_rank + self.qk_rope_head_dim, + config.kv_lora_rank + self.qk_rope_head_dim + + type(self).mla_head_padding, scale=self.softmax_scale, num_kv_heads=num_key_value_heads, v_head_size=config.kv_lora_rank, num_replicate_kv_heads=num_replicate_kv_heads, use_flash_mla=use_flash_mla, - mla_index_topk=config.index_topk) + mla_index_topk=(config.index_topk + if type(self).use_sparse_mla else None)) self.vc = DeepseekV2BMM(self.num_heads, config.kv_lora_rank, self.v_head_dim, dtype=dtype, device=device) self.o_proj = build_o_proj( diff --git a/lmdeploy/pytorch/models/glm4_1v.py b/lmdeploy/pytorch/models/glm4_1v.py index ac802fe6a2..eaa04f0517 100644 --- a/lmdeploy/pytorch/models/glm4_1v.py +++ b/lmdeploy/pytorch/models/glm4_1v.py @@ -158,8 +158,11 @@ class Glm4vVisionRotaryEmbedding(nn.Module): def __init__(self, dim: int, theta: float = 10000.0, device: torch.device = None) -> None: super().__init__() - inv_freq = 1.0 / (theta**(torch.arange(0, dim, 2, dtype=torch.int64).float() / dim)) - inv_freq = inv_freq.to(device=device) + inv_freq = 1.0 / (theta**(torch.arange(0, + dim, + 2, + dtype=torch.float32, + device=device) / dim)) self.register_buffer('inv_freq', inv_freq, persistent=False) def forward(self, seqlen: int) -> torch.Tensor: diff --git a/lmdeploy/pytorch/models/glm5_next.py b/lmdeploy/pytorch/models/glm5_next.py new file mode 100644 index 0000000000..17e46a388c --- /dev/null +++ b/lmdeploy/pytorch/models/glm5_next.py @@ -0,0 +1,1856 @@ +# Copyright (c) OpenMMLab. All rights reserved. +"""PyTorch engine implementation of multimodal GLM-5.3-Flash.""" + +from __future__ import annotations + +import re +from collections.abc import Iterable, Sequence +from functools import partial +from typing import Any + +import torch +import torch.nn.functional as F +from torch import distributed as dist +from torch import nn + +from lmdeploy.pytorch.backends.cuda.attention.tilelang_sparse_mla import ( + TilelangSparseMLADecode, +) +from lmdeploy.pytorch.backends.cuda.kpool import ( + kpool_compress_quantize_cuda, + kpool_score_contiguous_cuda, + kpool_score_paged_cuda, + kpool_select_groups_cuda, +) +from lmdeploy.pytorch.configurations.glm5_next import is_glm5_kda_layer +from lmdeploy.pytorch.consts import ( + GLM5_KDA_CONV_STATE, + GLM5_KDA_RECURRENT_STATE, + GLM5_KPOOL_TAIL_K_STATE, + GLM5_KPOOL_TAIL_SCORE_STATE, +) +from lmdeploy.pytorch.distributed import get_dist_manager, get_tp_world_rank +from lmdeploy.pytorch.model_inputs import StepContext, StepContextManager +from lmdeploy.pytorch.nn import ( + FlashAttention, + FP32LayerNorm, + HcPrePost, + Kda, + KPoolIndexer, + ParallelLMHead, + RMSNorm, + apply_rotary_pos_emb_fp32, +) +from lmdeploy.pytorch.nn.gated_delta import GatedDeltaMeta, build_rmsnorm_gated +from lmdeploy.pytorch.nn.kpool import ( + kpool_decode_update, + kpool_expand_selected_groups, + kpool_partition_update, + kpool_pooled_block_offsets, + kpool_read_packed_cache, + kpool_rotate_query, + kpool_write_packed_cache, + kpool_write_packed_cache_batched, +) +from lmdeploy.pytorch.nn.linear import ( + build_colwise_linear, + build_merged_colwise_linear, + build_o_proj, + build_qkv_proj, + build_rowwise_linear, +) +from lmdeploy.pytorch.nn.nsa import get_dsa_indexer_k_cache +from lmdeploy.pytorch.weight_loader.model_weight_loader import load_weight +from lmdeploy.vl.constants import Modality + +from .deepseek_v2 import DeepseekV2MLP, DeepseekV2MoE +from .deepseek_v32 import DeepseekV32Attention, DeepseekV32ForCausalLM +from .glm4_1v import ( + Glm4vVisionAttention, + Glm4vVisionPatchEmbed, + Glm4vVisionRotaryEmbedding, +) +from .qwen3_vl import Qwen3VLInputProcessor +from .utils.model import build_embedding, vlm_model + +Glm5NextVisionRMSNorm = RMSNorm + + +class Glm5NextLayerNorm(FP32LayerNorm): + """GLM-5.3 FP32 LayerNorm using LMDeploy's reusable implementation.""" + + +# Backward-compatible name retained for the vision numerical contract tests +# and downstream imports. The provider is also shared by the KPool indexer. +Glm5NextVisionLayerNorm = Glm5NextLayerNorm + + +def _build_glm53_latent_norm(hidden_size: int, eps: float, + dtype: torch.dtype | None, + device: torch.device | None) -> RMSNorm: + """Build the unquantized BF16 latent norm used by GLM sparse attention.""" + if dtype is None: + dtype = torch.get_default_dtype() + return RMSNorm(hidden_size, + eps, + quant_config=None, + dtype=dtype, + device=device) + + +def _glm_swiglu_impl(intermediate: torch.Tensor, + swiglu_limit: float, + precise_mul: bool = False) -> torch.Tensor: + """GLM/DeepSeek-V4 clamped SwiGLU used by dense and routed experts.""" + from lmdeploy.pytorch.kernels.cuda.activation import silu_and_mul + input_shape = intermediate.shape + intermediate = intermediate.flatten(0, -2) + output = silu_and_mul(intermediate, + swiglu_limit=swiglu_limit, + precise_mul=precise_mul) + return output.unflatten(0, input_shape[:-1]) + + +_GLM53_COMPACT_FP8_MOE_ACT = partial( + _glm_swiglu_impl, swiglu_limit=10.0, precise_mul=True) + + +class Glm5NextVisionPatchEmbed(Glm4vVisionPatchEmbed): + """Patch embedding with SGLang's unfold-plus-linear reduction order.""" + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + target_dtype = self.proj.weight.dtype + hidden_states = hidden_states.view( + -1, + self.in_channels, + self.temporal_patch_size, + self.patch_size, + self.patch_size, + ).to(dtype=target_dtype) + hidden_states = hidden_states.flatten(1) + weight = self.proj.weight.flatten(1) + return F.linear(hidden_states, weight, self.proj.bias) + + +class Glm5NextVisionAttention(Glm4vVisionAttention): + """GLM-OCR attention with GLM-5.3's per-head Q/K RMSNorm.""" + + # SGLang's GLM-5.3 block leaves VisionAttention.layer_norm_eps at its + # default instead of forwarding vision_config.rms_norm_eps. Keep this + # explicit because the outer block norms do use the config value. + qk_norm_eps = 1e-6 + + def __init__(self, + config: Any, + dtype: torch.dtype | None = None, + device: torch.device | None = None): + # Reuse LMDeploy's QKV/row-parallel projections and rotary operator. + # Vision weights stay in BF16 even when the language tower is block-FP8. + super().__init__(config, dtype=dtype, device=device) + self.q_norm = Glm5NextVisionRMSNorm(self.head_dim, + eps=self.qk_norm_eps, + quant_config=None, + dtype=dtype, + device=device) + self.k_norm = Glm5NextVisionRMSNorm(self.head_dim, + eps=self.qk_norm_eps, + quant_config=None, + dtype=dtype, + device=device) + + def forward( + self, + hidden_states: torch.Tensor, + cu_seqlens: torch.Tensor, + max_seqlen: int, + rotary_pos_emb: tuple[torch.Tensor, torch.Tensor], + ) -> torch.Tensor: + seq_length = hidden_states.shape[0] + qkv_states = self.qkv(hidden_states).flatten(0, -2) + query, key, value = self.qkv.split_qkv(qkv_states) + + query = self.q_norm(query) + key = self.k_norm(key) + cos, sin = rotary_pos_emb + query, key = apply_rotary_pos_emb_fp32(query, key, cos, sin) + output = self.attention( + query, + key, + value, + q_start_loc=cu_seqlens[:-1], + q_seqlens=cu_seqlens[1:] - cu_seqlens[:-1], + max_q_seqlen=max_seqlen, + ) + return self.proj(output.reshape(seq_length, -1)) + + +class Glm5NextVisionMLP(nn.Module): + """TP-aware vision MLP with GLM-5.3's clamped SwiGLU.""" + + def __init__(self, + config: Any, + dtype: torch.dtype | None = None, + device: torch.device | None = None): + super().__init__() + self.swiglu_limit = config.swiglu_limit + self.gate_up_proj = build_merged_colwise_linear( + in_features=config.hidden_size, + all_out_features=[config.intermediate_size, + config.intermediate_size], + bias=True, + dtype=dtype, + device=device, + quant_config=None, + is_tp=True, + ) + self.down_proj = build_rowwise_linear( + in_features=config.intermediate_size, + out_features=config.hidden_size, + bias=True, + dtype=dtype, + device=device, + quant_config=None, + is_tp=True, + ) + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + hidden_states = self.gate_up_proj(hidden_states) + hidden_states = _glm_swiglu_impl(hidden_states, + self.swiglu_limit, + precise_mul=True) + return self.down_proj(hidden_states) + + +class Glm5NextVisionPatchMerger(nn.Module): + """GLM-5.3 vision projector after spatial downsampling.""" + + def __init__(self, + config: Any, + dtype: torch.dtype | None = None, + device: torch.device | None = None): + super().__init__() + dim = config.out_hidden_size + context_dim = getattr(config, 'projection_intermediate_size', + config.intermediate_size) + self.swiglu_limit = config.swiglu_limit + self.proj = nn.Linear(dim, + dim, + bias=False, + dtype=dtype, + device=device) + self.post_projection_norm = Glm5NextLayerNorm(dim, + eps=1e-6, + device=device) + self.gate_up_proj = build_merged_colwise_linear( + in_features=dim, + all_out_features=[context_dim, context_dim], + bias=False, + dtype=dtype, + device=device, + quant_config=None, + is_tp=True, + ) + self.down_proj = build_rowwise_linear( + in_features=context_dim, + out_features=dim, + bias=False, + dtype=dtype, + device=device, + quant_config=None, + is_tp=True, + ) + self.act1 = nn.GELU() + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + hidden_states = self.proj(hidden_states) + hidden_states = self.act1(self.post_projection_norm(hidden_states)) + hidden_states = self.gate_up_proj(hidden_states) + hidden_states = _glm_swiglu_impl(hidden_states, + self.swiglu_limit, + precise_mul=True) + return self.down_proj(hidden_states) + + +class Glm5NextVisionBlock(nn.Module): + """A GLM-5.3 vision block built from public LMDeploy operators.""" + + def __init__(self, + config: Any, + dtype: torch.dtype | None = None, + device: torch.device | None = None): + super().__init__() + self.norm1 = Glm5NextVisionRMSNorm(config.hidden_size, + eps=config.rms_norm_eps, + quant_config=None, + dtype=dtype, + device=device) + self.norm2 = Glm5NextVisionRMSNorm(config.hidden_size, + eps=config.rms_norm_eps, + quant_config=None, + dtype=dtype, + device=device) + self.attn = Glm5NextVisionAttention(config, + dtype=dtype, + device=device) + self.mlp = Glm5NextVisionMLP(config, + dtype=dtype, + device=device) + + def forward( + self, + hidden_states: torch.Tensor, + cu_seqlens: torch.Tensor, + max_seqlen: int, + rotary_pos_emb: tuple[torch.Tensor, torch.Tensor], + ) -> torch.Tensor: + residual = hidden_states + hidden_states = self.norm1(hidden_states) + hidden_states = self.attn(hidden_states, + cu_seqlens=cu_seqlens, + max_seqlen=max_seqlen, + rotary_pos_emb=rotary_pos_emb) + hidden_states, residual = self.norm2(hidden_states, residual) + hidden_states = self.mlp(hidden_states) + return residual + hidden_states + + +@vlm_model +class Glm5NextVisionModel(nn.Module): + """Native GLM-5.3 image/video encoder.""" + + def __init__(self, + config: Any, + dtype: torch.dtype | None = None, + device: torch.device | None = None): + super().__init__() + self.config = config + self.spatial_merge_size = config.spatial_merge_size + self.patch_embed = Glm5NextVisionPatchEmbed(config, + dtype=dtype, + device=device) + head_dim = config.hidden_size // config.num_heads + self.rotary_pos_emb = Glm4vVisionRotaryEmbedding( + head_dim // 2, device=device) + self.blocks = nn.ModuleList([ + Glm5NextVisionBlock(config, dtype=dtype, device=device) + for _ in range(config.depth) + ]) + self.post_layernorm = Glm5NextVisionRMSNorm( + config.hidden_size, + eps=config.rms_norm_eps, + quant_config=None, + dtype=dtype, + device=device) + self.downsample = nn.Conv2d( + in_channels=config.hidden_size, + out_channels=config.out_hidden_size, + kernel_size=config.spatial_merge_size, + stride=config.spatial_merge_size, + dtype=dtype, + device=device, + ) + self.merger = Glm5NextVisionPatchMerger(config, + dtype=dtype, + device=device) + + def _rotary_embedding( + self, + grid_thw: torch.Tensor, + ) -> tuple[torch.Tensor, torch.Tensor]: + pos_ids = [] + merge = self.spatial_merge_size + for t, h, w in grid_thw.tolist(): + if h % merge or w % merge: + raise ValueError( + f'vision grid {(t, h, w)} is not divisible by merge={merge}.') + hpos = torch.arange(h).unsqueeze(1).expand(-1, w) + wpos = torch.arange(w).unsqueeze(0).expand(h, -1) + hpos = hpos.reshape(h // merge, merge, w // merge, + merge).permute(0, 2, 1, 3).flatten() + wpos = wpos.reshape(h // merge, merge, w // merge, + merge).permute(0, 2, 1, 3).flatten() + pos_ids.append(torch.stack([hpos, wpos], dim=-1).repeat(t, 1)) + pos_ids = torch.cat(pos_ids, dim=0) + max_grid_size = int(grid_thw[:, 1:].max().item()) + full = self.rotary_pos_emb(max_grid_size) + rotary = full[pos_ids.to(full.device)].flatten(1).repeat(1, 2) + return rotary.cos(), rotary.sin() + + def forward(self, pixel_values: torch.Tensor, + grid_thw: torch.Tensor) -> torch.Tensor: + expected_patches = int(grid_thw.prod(dim=-1).sum().item()) + hidden_states = self.patch_embed(pixel_values) + if hidden_states.shape[0] != expected_patches: + raise ValueError( + 'vision patch count mismatch: ' + f'got {hidden_states.shape[0]}, expected {expected_patches}.') + + lengths = torch.repeat_interleave(grid_thw[:, 1] * grid_thw[:, 2], + grid_thw[:, 0]) + max_seqlen = int(lengths.max().item()) + cu_seqlens = F.pad(lengths.cumsum(0, dtype=torch.int32), + (1, 0), + value=0).to(hidden_states.device) + rotary_pos_emb = self._rotary_embedding(grid_thw) + # Keep the trigonometric tables in FP32. The vision RoPE helper casts + # q/k to FP32 for the rotation and rounds only the final result. + rotary_pos_emb = tuple(x.to(device=hidden_states.device) + for x in rotary_pos_emb) + for block in self.blocks: + hidden_states = block(hidden_states, + cu_seqlens=cu_seqlens, + max_seqlen=max_seqlen, + rotary_pos_emb=rotary_pos_emb) + + hidden_states = self.post_layernorm(hidden_states) + merge = self.spatial_merge_size + hidden_states = hidden_states.view(-1, merge, merge, + hidden_states.shape[-1]) + hidden_states = hidden_states.permute(0, 3, 1, 2) + hidden_states = self.downsample(hidden_states) + hidden_states = hidden_states.reshape(-1, + self.config.out_hidden_size) + hidden_states = self.merger(hidden_states) + expected_tokens = expected_patches // (merge * merge) + if hidden_states.shape[0] != expected_tokens: + raise ValueError( + 'vision embedding count mismatch: ' + f'got {hidden_states.shape[0]}, expected {expected_tokens}.') + return hidden_states + + +class Glm5NextInputProcessor(Qwen3VLInputProcessor): + """Reuse the common image/video ``MultiModalData`` ownership boundary.""" + + def __init__(self, config: Any, dtype: torch.dtype | None) -> None: + super().__init__(config=config, dtype=dtype) + if config.vision_config.spatial_merge_size != 2: + raise ValueError('GLM-5.3 mRoPE currently requires merge size 2.') + + +class Glm5NextMLP(DeepseekV2MLP): + """DeepSeek MLP projections with GLM-5.3's activation clamp.""" + + def __init__(self, config: Any, *args, **kwargs): + super().__init__(config, *args, **kwargs) + self.swiglu_limit = config.swiglu_limit + + def forward(self, x: torch.Tensor) -> torch.Tensor: + gate_up = self.gate_up_proj(x) + return self.down_proj( + _glm_swiglu_impl(gate_up, + self.swiglu_limit, + precise_mul=True)) + + +class Glm5NextNoauxTCRouter(nn.Module): + """GLM noaux router using LMDeploy's existing Triton implementation.""" + + def __init__(self, config: Any): + super().__init__() + self.top_k = config.num_experts_per_tok + self.num_experts = config.n_routed_experts + self.n_group = config.n_group + self.topk_group = config.topk_group + self.scoring_func = config.scoring_func + self.renormalize = bool(config.norm_topk_prob and self.top_k > 1) + # Keep the generic router output normalized. GLM's model-level 2.5 + # factor is owned by ``fused_moe_output_scale`` below so it is applied + # once, after the FP32 routed-expert reduction. + self.routed_scaling_factor = 1.0 + self.router_n_groups = getattr(config, 'router_n_groups', -1) + contract = ( + getattr(config, 'model_type', None) == 'glm5_next_text' + and config.topk_method == 'noaux_tc' + and self.top_k == 8 + and self.num_experts == 288 + and self.n_group == 1 + and self.topk_group == 1 + and self.scoring_func == 'sigmoid' + and self.renormalize + and config.routed_scaling_factor == 2.5 + and self.router_n_groups == -1 + and getattr(config, 'moe_router_dtype', None) == 'float32' + and getattr(config, 'n_shared_experts', None) == 1 + ) + if not contract: + raise ValueError( + 'The shared GLM-5.3 router requires the exact official ' + 'glm5_next_text noaux/FP32/288-expert/topk8 contract.') + + def forward( + self, + router_logits: torch.Tensor, + correction_bias: torch.Tensor, + ) -> tuple[torch.Tensor, torch.Tensor]: + """Return normalized weights and expert ids.""" + from lmdeploy.pytorch.kernels.cuda.moe.route_noaux_tc import ( + fused_noaux_tc_routing, + ) + + return fused_noaux_tc_routing( + logits=router_logits, + bias=correction_bias, + top_k=self.top_k, + num_experts=self.num_experts, + n_group=self.n_group, + topk_group=self.topk_group, + renormalize=self.renormalize, + routed_scaling_factor=self.routed_scaling_factor, + ) + + +class Glm5NextMoE(DeepseekV2MoE): + """DeepSeek routed experts with the same clamp as GLM-5.3.""" + + # Keep only the model-specific clamp; dispatch, quantization and grouped + # GEMMs are owned by LMDeploy's generic blocked-FP8 MoE implementation. + fused_moe_act_func = staticmethod(_GLM53_COMPACT_FP8_MOE_ACT) + fused_moe_fp32_acc = True + # Match the GLM-5.3 contract: routing returns normalized, unscaled weights; + # the 2.5 routed scale is applied once to the FP32 expert reduction before + # its BF16 store. + router_routed_scaling_factor = 1.0 + fused_moe_output_scale = 2.5 + # Reuse LMDeploy's generic compact DeepGEMM path. With the model's + # sigmoid KDA gate and single-owner routed scaling restored, its real + # weight replay is closer to the reference than the generic Triton path. + fused_moe_use_deep_gemm = True + shared_expert_cls = Glm5NextMLP + + def __init__(self, config: Any, *args, **kwargs): + super().__init__(config, *args, **kwargs) + if self.gate.fake_eplb or self.gate.eplb_dispatch_info is not None: + raise RuntimeError( + 'The GLM-5.3 router does not permit fake ' + 'routing or EPLB expert remapping.') + self.gate.noaux_tc_router = Glm5NextNoauxTCRouter(config) + + def forward( + self, + hidden_states: torch.Tensor, + all_routed_experts: torch.Tensor | None = None, + ) -> torch.Tensor: + if all_routed_experts is not None: + raise RuntimeError( + 'GLM-5.3 routed-expert capture is not supported.') + return super().forward(hidden_states, all_routed_experts=None) + + +def _load_vector_shard(param: nn.Parameter, + loaded_weight: torch.Tensor) -> None: + """Load a flattened attention-head vector for the local TP rank.""" + world_size, rank = get_tp_world_rank('attn') + loaded_weight = loaded_weight.flatten() + if world_size > 1: + loaded_weight = loaded_weight.chunk(world_size, dim=0)[rank] + param.data.copy_(loaded_weight.to(device=param.device, dtype=param.dtype)) + + +class Glm5NextQKVConv1d(nn.Module): + """FP32 depthwise-convolution weight container for fused Q/K/V KDA.""" + + _SHARD_IDS = {'q': 0, 'k': 1, 'v': 2, 0: 0, 1: 1, 2: 2} + + def __init__(self, + local_projection_size: int, + kernel_size: int, + device: torch.device | None = None): + super().__init__() + self.local_projection_size = local_projection_size + weight = torch.empty(3 * local_projection_size, + 1, + kernel_size, + dtype=torch.float32, + device=device) + self.weight = nn.Parameter(weight) + self.weight.weight_loader = self._weight_loader + + def _weight_loader(self, param: nn.Parameter, loaded_weight: torch.Tensor, + shard_id: str | int) -> None: + world_size, rank = get_tp_world_rank('attn') + if world_size > 1: + loaded_weight = loaded_weight.chunk(world_size, dim=0)[rank] + shard = self._SHARD_IDS[shard_id] + start = shard * self.local_projection_size + target = param.data.narrow(0, start, self.local_projection_size) + target.copy_(loaded_weight.to(device=target.device, + dtype=target.dtype)) + + +class Glm5NextLinearAttention(nn.Module): + """GLM-5.3 Kimi Delta Attention projections and backend dispatch.""" + + def __init__(self, + config: Any, + layer_idx: int, + dtype: torch.dtype | None = None, + device: torch.device | None = None, + all_reduce: bool = True): + super().__init__() + linear_config = config.linear_attn_config + self.layer_idx = layer_idx + self.hidden_size = config.hidden_size + self.num_heads = linear_config['num_heads'] + self.head_dim = linear_config['head_dim'] + self.conv_kernel_size = linear_config['short_conv_kernel_size'] + self.lower_bound = linear_config.get('gate_lower_bound', -5.0) + + tp, _ = get_tp_world_rank('attn') + if self.num_heads % tp: + raise ValueError( + f'KDA heads={self.num_heads} is not divisible by attention TP={tp}.' + ) + self.local_num_heads = self.num_heads // tp + projection_size = self.num_heads * self.head_dim + local_projection_size = self.local_num_heads * self.head_dim + + # The official checkpoint keeps every KDA projection in BF16 even + # though the MLA/MLP weights use block FP8. Do not inherit the global + # quantization policy for these modules. + self.qkv_proj = build_qkv_proj( + self.hidden_size, + num_q_heads=self.num_heads, + num_kv_heads=self.num_heads, + head_size=self.head_dim, + bias=False, + quant_config=None, + dtype=dtype, + device=device, + is_tp=True, + ) + self.b_proj = build_colwise_linear( + self.hidden_size, + self.num_heads, + bias=False, + quant_config=None, + dtype=dtype, + device=device, + is_tp=True, + ) + self.f_a_proj = build_colwise_linear( + self.hidden_size, + self.head_dim, + bias=False, + quant_config=None, + dtype=dtype, + device=device, + is_tp=False, + ) + self.f_b_proj = build_colwise_linear( + self.head_dim, + projection_size, + bias=False, + quant_config=None, + dtype=dtype, + device=device, + is_tp=True, + ) + self.g_b_proj = build_colwise_linear( + self.head_dim, + projection_size, + bias=False, + quant_config=None, + dtype=dtype, + device=device, + is_tp=True, + ) + self.g_a_proj = build_colwise_linear( + self.hidden_size, + self.head_dim, + bias=False, + quant_config=None, + dtype=dtype, + device=device, + is_tp=False, + ) + + self.qkv_conv1d = Glm5NextQKVConv1d(local_projection_size, + self.conv_kernel_size, + device=device) + self.A_log = nn.Parameter( + torch.empty(self.local_num_heads, + dtype=torch.float32, + device=device)) + self.dt_bias = nn.Parameter( + torch.empty(local_projection_size, + dtype=torch.float32, + device=device)) + self.A_log.weight_loader = _load_vector_shard + self.dt_bias.weight_loader = _load_vector_shard + + self.o_norm = build_rmsnorm_gated( + self.head_dim, + eps=config.rms_norm_eps, + activation='sigmoid', + dtype=dtype, + device=device, + ) + self.o_proj = build_o_proj( + projection_size, + self.hidden_size, + bias=False, + quant_config=None, + dtype=dtype, + device=device, + is_tp=True, + all_reduce=all_reduce, + ) + self.kda = Kda() + + def forward(self, hidden_states: torch.Tensor, + past_key_value: Sequence[torch.Tensor], + kda_metadata: GatedDeltaMeta) -> torch.Tensor: + mixed_qkv = self.qkv_proj(hidden_states) + raw_beta = self.b_proj(hidden_states) + raw_gate = self.f_b_proj(self.f_a_proj(hidden_states)) + norm_gate = self.g_b_proj(self.g_a_proj(hidden_states)) + + core_output = self.kda( + mixed_qkv=mixed_qkv, + raw_gate=raw_gate, + raw_beta=raw_beta, + conv_weight=self.qkv_conv1d.weight, + conv_bias=None, + a_log=self.A_log, + dt_bias=self.dt_bias, + conv_state=past_key_value[0], + recurrent_state=past_key_value[1], + metadata=kda_metadata, + num_heads=self.local_num_heads, + head_dim=self.head_dim, + lower_bound=self.lower_bound, + ) + norm_gate = norm_gate.unflatten(-1, + (self.local_num_heads, self.head_dim)) + output_shape = core_output.shape + core_output = self.o_norm(core_output.reshape(-1, self.head_dim), + norm_gate.reshape(-1, self.head_dim)) + core_output = core_output.view(output_shape).flatten(-2, -1) + return self.o_proj(core_output) + + +class Glm5NextSparseAttention(DeepseekV32Attention): + """GLM MLA without RoPE and with a pageable KPool-4 indexer.""" + + use_sparse_mla = False + mla_head_padding = 64 + + def __init__(self, + config: Any, + layer_idx: int, + dtype: torch.dtype | None = None, + device: torch.device | None = None, + all_reduce: bool = True): + super().__init__(config, + layer_idx, + dtype=dtype, + device=device, + all_reduce=all_reduce) + # DeepSeek keeps these latent-norm parameters in FP32. GLM-5.3's + # checkpoint and SGLang runtime keep them in the activation dtype; + # rebuild only these two containers before weight loading. + if self.q_lora_rank is not None: + self.q_a_layernorm = _build_glm53_latent_norm( + config.q_lora_rank, config.rms_norm_eps, dtype, device) + self.kv_a_layernorm = _build_glm53_latent_norm( + config.kv_lora_rank, config.rms_norm_eps, dtype, device) + self.index_topk = config.index_topk + self.index_kpool = config.index_kpool + try: + self.cache_layer_idx = config.full_attention_layer_ids.index( + layer_idx) + except ValueError as error: + raise ValueError( + f'GLM-5.3 full-attention layer {layer_idx} is missing from ' + 'the compact cache map.') from error + + # Keep the checkpoint's BF16 KV-B projection alongside the absorbed + # KC/VC views. Short prefill uses the former to reproduce SGLang's + # decompressed dense MHA; decode continues to use KC/VC. + self.kv_b_proj = build_colwise_linear( + self.kv_lora_rank, + self.num_heads * (self.qk_nope_head_dim + self.v_head_dim), + bias=False, + dtype=dtype, + device=device, + is_tp=True, + quant_config=None, + dp_disable_tp=True, + ) + self.prefill_attn_fwd = FlashAttention( + self.num_heads, + self.q_head_dim, + scale=self.softmax_scale, + v_head_dim=self.v_head_dim, + causal=True, + ) + self.decode_attn_fwd = TilelangSparseMLADecode( + index_topk=self.index_topk, + index_kpool=getattr(config, 'index_kpool', 1), + ) + + def _build_indexer(self, config: Any, layer_idx: int, dtype: torch.dtype, + device: torch.device): + del layer_idx + return KPoolIndexer( + hidden_size=config.hidden_size, + index_n_heads=config.index_n_heads, + index_head_dim=config.index_head_dim, + index_topk=config.index_topk, + q_lora_rank=config.q_lora_rank, + index_kpool=config.index_kpool, + dtype=dtype, + device=device, + key_norm=Glm5NextLayerNorm(config.index_head_dim, + eps=1e-6, + device=device), + ) + + def _qkv_proj_unabsorbed(self, hidden_states: torch.Tensor, + num_heads: int): + """Project raw per-head Q and latent KV without absorption.""" + nope_size = self.kv_lora_rank + pe_size = self.qk_rope_head_dim + if self.q_lora_rank is None: + q_a_states = hidden_states + key_states = self.kv_a_proj_with_mqa(hidden_states[0, :, None]) + else: + q_a_states, key_states = self.fused_qkv_a_proj( + hidden_states).split([self.q_lora_rank, nope_size + pe_size], + dim=-1) + key_states = key_states[0, :, None] + + q_len = q_a_states.size(1) + if self.q_lora_rank is None: + q_lora = q_a_states + query = self.q_proj(q_a_states) + else: + q_lora = self.q_a_layernorm(q_a_states) + query = self.q_b_proj(q_lora) + query = query.view(q_len, num_heads, self.q_head_dim) + + key_states, value_states, _ = self._kv_proj(key_states, nope_size) + return query, key_states, value_states, q_lora + + def _kv_proj(self, key_states: torch.Tensor, nope_size: int): + """Normalize latent KV with GLM-5.3's exact FP32 math order.""" + k_pe = key_states[..., nope_size:] + value_states = key_states[..., :nope_size] + value_states = self.kv_a_layernorm(value_states) + key_states[..., :nope_size] = value_states + return key_states, value_states, k_pe + + def _update_kpool_cache( + self, + hidden_states: torch.Tensor, + tail_state: Sequence[torch.Tensor], + state_ids: torch.Tensor, + attn_metadata: Any, + ) -> torch.Tensor: + """Compress closed pools and persist each request's unfinished tail.""" + if tail_state is None or len(tail_state) != 2: + raise RuntimeError( + 'GLM-5.3 KPool requires key and score tail state caches.') + if state_ids is None: + raise RuntimeError('GLM-5.3 KPool requires stable state cache ids.') + + tail_k_state, tail_score_state = tail_state + indexer_k_cache = get_dsa_indexer_k_cache(self.cache_layer_idx) + key = self.indexer.project_key(hidden_states)[0] + score = self.indexer.project_compress_score(hidden_states)[0] + if attn_metadata.is_decoding: + if key.size(0) != state_ids.numel(): + raise RuntimeError( + 'KPool decode requires one token per request state.') + history_lengths = (attn_metadata.kv_seqlens + - attn_metadata.q_seqlens) + update = kpool_decode_update( + key, + score, + tail_k_state, + tail_score_state, + state_ids, + history_lengths, + self.index_kpool, + ) + pooled_fp8, pooled_scale = kpool_compress_quantize_cuda( + update.closed_keys, + update.closed_scores, + self.indexer.index_kpool_compress_ape, + mode='decode', + round_scale=self.indexer.scale_fmt is not None, + ) + kpool_write_packed_cache_batched( + indexer_k_cache, + attn_metadata.block_offsets, + update.group_ids, + pooled_fp8, + pooled_scale, + self.index_kpool, + update.should_close, + ) + tail_k_state.index_copy_(0, update.safe_state_ids, + update.next_tail_keys) + tail_score_state.index_copy_(0, update.safe_state_ids, + update.next_tail_scores) + return indexer_k_cache + + q_seqlens = attn_metadata.q_seqlens.tolist() + kv_seqlens = attn_metadata.kv_seqlens.tolist() + if len(q_seqlens) != len(kv_seqlens): + raise RuntimeError('KPool query/KV batch lengths do not match.') + if state_ids.numel() != len(q_seqlens): + raise RuntimeError( + 'KPool state id count does not match the request batch.') + + token_offset = 0 + for batch_idx, (q_len, kv_len) in enumerate( + zip(q_seqlens, kv_seqlens)): + q_len = int(q_len) + kv_len = int(kv_len) + history_len = kv_len - q_len + if history_len < 0: + raise RuntimeError( + f'KPool received q_len={q_len} greater than kv_len={kv_len}.') + state_id = int(state_ids[batch_idx].item()) + previous_tail_len = history_len % self.index_kpool + if state_id >= 0: + previous_tail_k = tail_k_state[ + state_id, :previous_tail_len] + previous_tail_score = tail_score_state[ + state_id, :previous_tail_len] + else: + previous_tail_k = key.new_zeros( + previous_tail_len, self.indexer.head_dim) + previous_tail_score = score.new_zeros( + previous_tail_len, self.indexer.head_dim) + + token_end = token_offset + q_len + update = kpool_partition_update( + key[token_offset:token_end], + score[token_offset:token_end], + history_length=history_len, + pool_size=self.index_kpool, + tail_keys=previous_tail_k, + tail_scores=previous_tail_score, + ) + if update.closed_group_ids.numel(): + pooled_fp8, pooled_scale = kpool_compress_quantize_cuda( + update.closed_keys, + update.closed_scores, + self.indexer.index_kpool_compress_ape, + mode=('decode' + if attn_metadata.is_decoding else 'extend'), + round_scale=self.indexer.scale_fmt is not None, + ) + kpool_write_packed_cache( + indexer_k_cache, + attn_metadata.block_offsets[batch_idx], + update.closed_group_ids, + pooled_fp8, + pooled_scale, + self.index_kpool, + ) + + if state_id >= 0: + tail_k_state[state_id].zero_() + tail_score_state[state_id].zero_() + tail_len = update.tail_keys.size(0) + if tail_len: + tail_k_state[state_id, :tail_len].copy_( + update.tail_keys) + tail_score_state[state_id, :tail_len].copy_( + update.tail_scores) + token_offset = token_end + + if token_offset != key.size(0): + raise RuntimeError( + f'KPool metadata accounts for {token_offset} tokens, ' + f'but projections contain {key.size(0)}.') + return indexer_k_cache + + def _select_kpool_indices( + self, + hidden_states: torch.Tensor, + q_lora: torch.Tensor, + indexer_k_cache: torch.Tensor, + attn_metadata: Any, + ) -> torch.Tensor: + """Score/select on attention-TP rank 0, then broadcast logical ids.""" + dist_ctx = get_dist_manager().current_context() + tp_group = dist_ctx.attn_tp_group + is_owner = tp_group.rank == 0 + total_rows = hidden_states.size(1) + output_width = self.index_topk + self.index_kpool - 1 + + if is_owner: + query = self.indexer.project_query(q_lora)[0] + query = kpool_rotate_query(query) + query_fp8, query_scale = self.indexer.quantize_fp8(query) + head_gate = self.indexer.project_head_gate(hidden_states)[0] + query_weight = (head_gate * query_scale.squeeze(-1) + * self.indexer.softmax_scale) + if attn_metadata.is_decoding: + seq_lens = attn_metadata.kv_seqlens.to(torch.int64) + group_lengths = torch.div( + seq_lens, + self.index_kpool, + rounding_mode='floor', + ) + pooled_block_offsets = kpool_pooled_block_offsets( + attn_metadata.block_offsets, + self.index_kpool, + ) + logits = kpool_score_paged_cuda( + query_fp8, + query_weight, + indexer_k_cache, + group_lengths, + pooled_block_offsets, + ) + selected_groups = kpool_select_groups_cuda( + logits.contiguous(), + group_lengths, + group_topk=self.index_topk // self.index_kpool, + ) + logical_indices = kpool_expand_selected_groups( + selected_groups, + group_lengths, + self.index_kpool, + self.index_topk, + seq_lens=seq_lens, + ) + else: + logical_indices = self._select_kpool_indices_prefill( + query_fp8, + query_weight, + indexer_k_cache, + attn_metadata, + ) + else: + logical_indices = torch.empty( + total_rows, + output_width, + dtype=torch.int32, + device=hidden_states.device, + ) + + if dist_ctx.dist_config.attn_tp > 1: + group = tp_group.gpu_group + source_rank = dist_ctx.rank - tp_group.rank + dist.broadcast(logical_indices, src=source_rank, group=group) + return logical_indices + + def _select_kpool_indices_prefill( + self, + query_fp8: torch.Tensor, + query_weight: torch.Tensor, + indexer_k_cache: torch.Tensor, + attn_metadata: Any, + ) -> torch.Tensor: + """Retain the ragged eager implementation for chunked prefill.""" + q_seqlens = attn_metadata.q_seqlens.tolist() + kv_seqlens = attn_metadata.kv_seqlens.tolist() + logical_parts = [] + token_offset = 0 + for batch_idx, (q_len, kv_len) in enumerate( + zip(q_seqlens, kv_seqlens)): + q_len = int(q_len) + kv_len = int(kv_len) + history_len = kv_len - q_len + num_groups = kv_len // self.index_kpool + token_end = token_offset + q_len + seq_lens = history_len + torch.arange( + 1, + q_len + 1, + dtype=torch.int64, + device=query_fp8.device, + ) + group_lengths = torch.div( + seq_lens, + self.index_kpool, + rounding_mode='floor', + ) + query_slice = query_fp8[token_offset:token_end] + weight_slice = query_weight[token_offset:token_end] + pooled_key, pooled_scale = kpool_read_packed_cache( + indexer_k_cache, + attn_metadata.block_offsets[batch_idx], + num_groups, + self.index_kpool, + ) + logits = kpool_score_contiguous_cuda( + query_slice, + weight_slice, + pooled_key, + pooled_scale, + group_lengths, + ) + group_budget = self.index_topk // self.index_kpool + selected_groups = kpool_select_groups_cuda( + logits.contiguous(), + group_lengths, + group_topk=group_budget, + max_group_length=num_groups, + ) + logical_parts.append( + kpool_expand_selected_groups( + selected_groups, + group_lengths, + self.index_kpool, + self.index_topk, + seq_lens=seq_lens, + )) + token_offset = token_end + return torch.cat(logical_parts, dim=0) + + def _kpool_indices( + self, + hidden_states: torch.Tensor, + q_lora: torch.Tensor, + tail_state: Sequence[torch.Tensor], + state_ids: torch.Tensor, + attn_metadata: Any, + return_indices: bool, + ) -> torch.Tensor | None: + indexer_k_cache = self._update_kpool_cache( + hidden_states, tail_state, state_ids, attn_metadata) + if not return_indices: + return None + return self._select_kpool_indices( + hidden_states, q_lora, indexer_k_cache, attn_metadata) + + def _absorbed_query(self, query: torch.Tensor, + num_heads: int) -> torch.Tensor: + query_states = query.new_empty( + query.size(0), num_heads, + self.kv_lora_rank + self.qk_rope_head_dim) + self.kc(query[..., :self.qk_nope_head_dim], + query_states[..., :self.kv_lora_rank]) + return query_states + + def _forward_prefill_mha( + self, + query: torch.Tensor, + key_states: torch.Tensor, + past_key_value: Sequence[torch.Tensor], + attn_metadata: Any, + num_heads: int, + ) -> torch.Tensor: + """Run short prefill as decompressed dense MHA over full latent KV.""" + k_scales_zeros = None if len( + past_key_value) == 2 else past_key_value[2] + v_scales_zeros = None if len( + past_key_value) == 2 else past_key_value[3] + flatten_latent = self.attn_fwd.fill_and_flatten_latent_kv_cache( + key_states, + past_key_value[0], + attn_metadata, + out_dtype=query.dtype, + k_scales_zeros=k_scales_zeros, + v_scales_zeros=v_scales_zeros, + ) + flatten_latent = flatten_latent[..., :self.kv_lora_rank].flatten(0, 1) + kv_states = self.kv_b_proj(flatten_latent) + kv_states = kv_states.view(-1, num_heads, + self.qk_nope_head_dim + self.v_head_dim) + key, value = kv_states.split([self.qk_nope_head_dim, self.v_head_dim], + dim=-1) + attn_output = self.prefill_attn_fwd( + query, + key, + value, + q_start_loc=attn_metadata.cu_seqlens_q[:-1], + q_seqlens=(attn_metadata.cu_seqlens_q[1:] + - attn_metadata.cu_seqlens_q[:-1]), + kv_start_loc=attn_metadata.cu_seqlens_k[:-1], + kv_seqlens=(attn_metadata.cu_seqlens_k[1:] + - attn_metadata.cu_seqlens_k[:-1]), + max_q_seqlen=attn_metadata.max_q_seqlen, + ) + return self.o_proj(attn_output.flatten(-2, -1)[None]) + + def forward( + self, + hidden_states: torch.Tensor, + past_key_value: Sequence[torch.Tensor], + attn_metadata: Any = None, + kpool_tail_state: Sequence[torch.Tensor] | None = None, + state_ids: torch.Tensor | None = None, + ) -> torch.Tensor: + dist_ctx = get_dist_manager().current_context() + num_heads = self.num_heads // dist_ctx.dist_config.attn_tp + nope_size = self.kv_lora_rank + q_len = hidden_states.size(1) + + (unabsorbed_query, key_states, value_states, + q_lora) = self._qkv_proj_unabsorbed( + hidden_states, num_heads=num_heads) + # The latent cache retains FlashMLA's 576-wide DeepSeek layout. GLM + # has no RoPE tail, so its final 64 dimensions are exact zeros. + key_states = F.pad(key_states, (0, self.mla_head_padding)) + + if not attn_metadata.is_decoding: + use_sparse = int(attn_metadata.max_kv_seqlen) > self.index_topk + logical_indices = self._kpool_indices( + hidden_states, + q_lora, + kpool_tail_state, + state_ids, + attn_metadata, + return_indices=use_sparse, + ) + if not use_sparse: + return self._forward_prefill_mha( + unabsorbed_query, + key_states, + past_key_value, + attn_metadata, + num_heads, + ) + + query_states = self._absorbed_query( + unabsorbed_query, num_heads) + attn_output = self.decode_attn_fwd.forward_prefill( + query_states, + key_states, + past_key_value[0], + attn_metadata, + scale=self.softmax_scale, + cache_writer=self.attn_fwd, + logical_indices=logical_indices, + k_scales_zeros=(None if len(past_key_value) == 2 else + past_key_value[2]), + v_scales_zeros=(None if len(past_key_value) == 2 else + past_key_value[3]), + ) + attn_bmm_out = attn_output.new_empty( + q_len, num_heads, self.v_head_dim) + self.vc(attn_output, attn_bmm_out) + return self.o_proj(attn_bmm_out.flatten(-2, -1)[None]) + + logical_indices = self._kpool_indices( + hidden_states, + q_lora, + kpool_tail_state, + state_ids, + attn_metadata, + return_indices=True, + ) + # GLM has no RoPE tail, so the absorbed query contains exactly 512 + # values; the cache retains its 576-wide FlashMLA storage alignment. + query_states = self._absorbed_query(unabsorbed_query, num_heads) + attn_output = self.decode_attn_fwd.forward( + query_states, + key_states, + value_states, + past_key_value[0], + past_key_value[0][..., :nope_size], + attn_metadata, + scale=self.softmax_scale, + cache_writer=self.attn_fwd, + k_scales_zeros=(None if len(past_key_value) == 2 else + past_key_value[2]), + v_scales_zeros=(None if len(past_key_value) == 2 else + past_key_value[3]), + logical_indices=logical_indices, + ) + attn_bmm_out = attn_output.new_empty(q_len, num_heads, self.v_head_dim) + self.vc(attn_output, attn_bmm_out) + return self.o_proj(attn_bmm_out.flatten(-2, -1)[None]) + + +class Glm5NextDecoderLayer(nn.Module): + """Hybrid KDA/MLA decoder block with mHC pre/post mixing.""" + + def __init__(self, + config: Any, + layer_idx: int, + dtype: torch.dtype | None = None, + device: torch.device | None = None): + super().__init__() + self.config = config + self.layer_idx = layer_idx + self.is_linear_attention = is_glm5_kda_layer(config, layer_idx) + if self.is_linear_attention: + self.self_attn = Glm5NextLinearAttention(config, + layer_idx, + dtype=dtype, + device=device) + else: + self.self_attn = Glm5NextSparseAttention(config, + layer_idx, + dtype=dtype, + device=device) + + mlp_layer_types = getattr(config, 'mlp_layer_types', None) + is_sparse = (mlp_layer_types is not None + and mlp_layer_types[layer_idx] == 'sparse') + if mlp_layer_types is None: + is_sparse = (config.n_routed_experts is not None + and layer_idx >= config.first_k_dense_replace + and layer_idx % config.moe_layer_freq == 0) + self.mlp = (Glm5NextMoE(config, layer_idx, dtype=dtype, device=device) + if is_sparse else Glm5NextMLP( + config, dtype=dtype, device=device)) + + self.input_layernorm = RMSNorm(config.hidden_size, + config.rms_norm_eps, + quant_config=None, + dtype=dtype, + device=device) + self.post_attention_layernorm = RMSNorm(config.hidden_size, + config.rms_norm_eps, + quant_config=None, + dtype=dtype, + device=device) + self.hc_prepost = HcPrePost(config.hc_mult, + config.hc_sinkhorn_iters, + config.hc_eps) + mix_hc = (2 + config.hc_mult) * config.hc_mult + hc_dim = config.hc_mult * config.hidden_size + self.hc_attn_fn = nn.Parameter(torch.empty(mix_hc, + hc_dim, + dtype=torch.float32, + device=device), + requires_grad=False) + self.hc_ffn_fn = nn.Parameter(torch.empty(mix_hc, + hc_dim, + dtype=torch.float32, + device=device), + requires_grad=False) + self.hc_attn_base = nn.Parameter(torch.empty(mix_hc, + dtype=torch.float32, + device=device), + requires_grad=False) + self.hc_ffn_base = nn.Parameter(torch.empty(mix_hc, + dtype=torch.float32, + device=device), + requires_grad=False) + self.hc_attn_scale = nn.Parameter(torch.empty(3, + dtype=torch.float32, + device=device), + requires_grad=False) + self.hc_ffn_scale = nn.Parameter(torch.empty(3, + dtype=torch.float32, + device=device), + requires_grad=False) + + def _hc_pre(self, hidden_states: torch.Tensor, fn: torch.Tensor, + scale: torch.Tensor, base: torch.Tensor, norm: RMSNorm): + return self.hc_prepost.pre( + hidden_states, fn, scale, base, norm.eps) + + def forward(self, hidden_states: torch.Tensor, + past_key_value: Sequence[torch.Tensor], attn_metadata: Any, + kda_metadata: GatedDeltaMeta, + kpool_tail_state: Sequence[torch.Tensor] | None = None, + state_ids: torch.Tensor | None = None) -> torch.Tensor: + residual = hidden_states + hidden_states, post, comb = self._hc_pre( + hidden_states, + self.hc_attn_fn, + self.hc_attn_scale, + self.hc_attn_base, + self.input_layernorm, + ) + hidden_states = self.input_layernorm(hidden_states) + if self.is_linear_attention: + hidden_states = self.self_attn(hidden_states, + past_key_value=past_key_value, + kda_metadata=kda_metadata) + else: + hidden_states = self.self_attn(hidden_states, + past_key_value=past_key_value, + attn_metadata=attn_metadata, + kpool_tail_state=kpool_tail_state, + state_ids=state_ids) + hidden_states = self.hc_prepost.post_expand(hidden_states, residual, + post, comb) + + residual = hidden_states + hidden_states, post, comb = self._hc_pre( + hidden_states, + self.hc_ffn_fn, + self.hc_ffn_scale, + self.hc_ffn_base, + self.post_attention_layernorm, + ) + hidden_states = self.post_attention_layernorm(hidden_states) + hidden_states = self.mlp(hidden_states) + return self.hc_prepost.post_expand(hidden_states, residual, post, comb) + + +class Glm5NextModel(nn.Module): + """GLM-5.3 hybrid backbone shared by text and multimodal generation.""" + + def __init__(self, + config: Any, + dtype: torch.dtype | None = None, + device: torch.device | None = None): + super().__init__() + self.config = config + self.embed_tokens = build_embedding(config.vocab_size, + config.hidden_size, + config.pad_token_id, + dtype=dtype, + device=device, + is_tp=True) + self.layers = nn.ModuleList([ + Glm5NextDecoderLayer(config, layer_idx, dtype=dtype, device=device) + for layer_idx in range(config.num_hidden_layers) + ]) + self.norm = RMSNorm(config.hidden_size, + config.rms_norm_eps, + quant_config=None, + dtype=dtype, + device=device) + + def forward( + self, + input_ids: torch.LongTensor, + position_ids: torch.LongTensor | None, + past_key_values: list[Sequence[torch.Tensor]], + attn_metadata: Any, + state_ids: torch.Tensor, + kpool_tail_states: list[Sequence[torch.Tensor]], + inputs_embeds: torch.Tensor | None = None, + ) -> torch.Tensor: + del position_ids # GLM-5.3 text attention has qk_rope_head_dim == 0. + if state_ids is None: + raise RuntimeError('GLM-5.3 KDA requires stable state cache ids.') + if kpool_tail_states is None: + raise RuntimeError('GLM-5.3 requires KPool tail state caches.') + if inputs_embeds is None: + inputs_embeds = self.embed_tokens(input_ids) + + hidden_states = inputs_embeds.unsqueeze(2).repeat( + 1, 1, self.config.hc_mult, 1) + kda_metadata = GatedDeltaMeta(hidden_states.size(1), + self.config.linear_conv_kernel_dim, + state_ids, attn_metadata) + if len(past_key_values) != len(self.layers): + raise RuntimeError( + f'GLM-5.3 expects {len(self.layers)} layer caches, got {len(past_key_values)}.' + ) + expected_full_layers = len(self.config.full_attention_layer_ids) + if len(kpool_tail_states) != expected_full_layers: + raise RuntimeError( + f'GLM-5.3 expects {expected_full_layers} KPool tail rows, ' + f'got {len(kpool_tail_states)}.') + full_layer_row = 0 + for layer, past_key_value in zip(self.layers, past_key_values): + kpool_tail_state = None + if not layer.is_linear_attention: + kpool_tail_state = kpool_tail_states[full_layer_row] + full_layer_row += 1 + hidden_states = layer(hidden_states, + past_key_value=past_key_value, + attn_metadata=attn_metadata, + kda_metadata=kda_metadata, + kpool_tail_state=kpool_tail_state, + state_ids=state_ids) + hidden_states = hidden_states.mean(dim=2) + return self.norm(hidden_states) + + def get_input_embeddings(self): + return self.embed_tokens + + +class Glm5NextForConditionalGeneration(DeepseekV32ForCausalLM): + """GLM-5.3 conditional-generation wrapper for text, image and video.""" + + def __init__(self, + config: Any, + ctx_mgr: StepContextManager, + dtype: torch.dtype | None = None, + device: torch.device | None = None): + nn.Module.__init__(self) + self.mm_config = config + self.config = config.text_config + self.quantization_config = getattr(config, 'quantization_config', None) + if self.quantization_config is not None: + self.config.quantization_config = self.quantization_config + self.dtype = dtype + self.ctx_mgr = ctx_mgr + self.input_processor = Glm5NextInputProcessor(config, dtype) + self.visual = Glm5NextVisionModel(config.vision_config, + dtype=dtype, + device=device) + self.model = Glm5NextModel(self.config, dtype=dtype, device=device) + self.lm_head = ParallelLMHead(self.config.vocab_size, + self.config.hidden_size, + bias=False, + dtype=dtype, + device=device) + if self.config.tie_word_embeddings: + self.lm_head.tie_weights(self.model.get_input_embeddings()) + self._load_buffers = {} + + def forward( + self, + input_ids: torch.Tensor, + position_ids: torch.Tensor, + past_key_values: list[Sequence[torch.Tensor]], + attn_metadata: Any = None, + inputs_embeds: torch.Tensor | None = None, + state_ids: torch.Tensor | None = None, + kpool_tail_states: list[Sequence[torch.Tensor]] | None = None, + pixel_values: torch.Tensor | None = None, + grid_thw: torch.Tensor | None = None, + vision_groups: list[dict[str, Any]] | None = None, + vision_prompt_order: list[int] | None = None, + multimodal_mask: torch.Tensor | None = None, + **kwargs, + ) -> torch.Tensor: + if inputs_embeds is None and (vision_groups or + pixel_values is not None): + inputs_embeds = self.get_input_embeddings()(input_ids) + if vision_groups: + grouped_outputs = [] + for group in vision_groups: + group_values = group['pixel_values'].to( + dtype=inputs_embeds.dtype) + grouped_outputs.append( + self.visual(group_values, group['grid_thw'])) + if vision_prompt_order is None: + if len(vision_groups) != 1: + raise ValueError( + 'vision_prompt_order is required for multiple ' + 'modality groups.') + vision_prompt_order = vision_groups[0]['flat_indices'] + vision_embeddings = self._restore_vision_prompt_order( + vision_groups, grouped_outputs, vision_prompt_order) + else: + # Keep the legacy single visual-call input contract available + # to callers that still pass pixel_values/grid_thw directly. + pixel_values = pixel_values.to(dtype=inputs_embeds.dtype) + vision_embeddings = self.visual(pixel_values, grid_thw) + num_slots = int(multimodal_mask.sum().item()) + if num_slots != vision_embeddings.shape[0]: + raise ValueError( + 'multimodal token/embedding count mismatch: ' + f'{num_slots} token slots for ' + f'{vision_embeddings.shape[0]} embeddings.') + scatter_mask = multimodal_mask.unsqueeze(-1).expand_as( + inputs_embeds) + inputs_embeds = inputs_embeds.masked_scatter( + scatter_mask, vision_embeddings.to(inputs_embeds)) + return self.model(input_ids=input_ids, + position_ids=position_ids, + past_key_values=past_key_values, + attn_metadata=attn_metadata, + inputs_embeds=inputs_embeds, + state_ids=state_ids, + kpool_tail_states=kpool_tail_states) + + def get_input_embeddings(self): + return self.model.get_input_embeddings() + + def get_input_processor(self): + """Return the model-specific image/video input processor.""" + return self.input_processor + + @staticmethod + def _multimodal_token_mask( + input_ids: torch.Tensor, + mm_inputs: list[Any], + ) -> torch.Tensor: + """Build a mask from the token ids owned by the MM input records.""" + token_ids = set() + for item in mm_inputs: + meta = item.meta or {} + token_id = (meta.get('image_token_id') + if item.modality == Modality.IMAGE else + meta.get('video_token_id')) + if token_id is not None: + token_ids.add(int(token_id)) + mask = torch.zeros_like(input_ids, dtype=torch.bool) + for token_id in token_ids: + mask |= input_ids == token_id + return mask + + @staticmethod + def _vision_grid(item: Any) -> list[torch.Tensor]: + """Return image grids, splitting video temporal units like SGLang.""" + grid = torch.as_tensor(item.meta['grid_thw'], + dtype=torch.long).reshape(3).cpu() + if item.modality != Modality.VIDEO: + return [grid] + t, h, w = grid.tolist() + return [torch.tensor([1, h, w], dtype=torch.long) for _ in range(t)] + + def _group_vision_inputs( + self, + input_multimodals: list[dict[str, Any]], + ) -> tuple[list[dict[str, Any]], list[int], list[Any]]: + """Pack visual calls by modality while retaining prompt item order.""" + records = [] + for batch_index, batch_item in enumerate(input_multimodals): + for item in batch_item.get('mm_data', []): + records.append( + dict(batch_index=batch_index, + flat_index=len(records), + item=item)) + + groups = [] + for modality in (Modality.IMAGE, Modality.VIDEO, Modality.AUDIO): + selected = [ + record for record in records + if record['item'].modality == modality + ] + if not selected: + continue + split_sizes = [ + int(record['item'].end - record['item'].start) + for record in selected + ] + if any(size <= 0 for size in split_sizes): + raise ValueError( + 'multimodal spans must have positive lengths.') + groups.append( + dict(modality=modality, + pixel_values=torch.cat( + [record['item'].data for record in selected], dim=0), + grid_thw=torch.stack([ + grid for record in selected + for grid in self._vision_grid(record['item']) + ], + dim=0), + flat_indices=[ + record['flat_index'] for record in selected + ], + split_sizes=split_sizes)) + + prompt_order = [ + record['flat_index'] + for record in sorted( + records, + key=lambda record: (record['batch_index'], + record['item'].start, + record['item'].end, + record['flat_index'])) + ] + return groups, prompt_order, [record['item'] for record in records] + + @staticmethod + def _restore_vision_prompt_order( + vision_groups: list[dict[str, Any]], + grouped_outputs: list[torch.Tensor], + prompt_order: list[int], + ) -> torch.Tensor: + """Split modality outputs per item and concatenate by prompt span.""" + if len(vision_groups) != len(grouped_outputs): + raise ValueError( + 'one visual output tensor is required per modality group.') + + chunks_by_flat_index = {} + for group, output in zip(vision_groups, grouped_outputs): + split_sizes = group['split_sizes'] + expected_rows = sum(split_sizes) + if output.shape[0] != expected_rows: + modality = group['modality'] + raise ValueError( + f'{modality.value} visual output has {output.shape[0]} ' + f'rows; expected {expected_rows} from prompt spans.') + chunks = torch.split(output, split_sizes, dim=0) + for flat_index, chunk in zip(group['flat_indices'], chunks): + if flat_index in chunks_by_flat_index: + raise ValueError( + f'duplicate visual flat index: {flat_index}.') + chunks_by_flat_index[flat_index] = chunk + + if set(chunks_by_flat_index) != set(prompt_order): + raise ValueError( + 'visual items and prompt-order items do not match.') + return torch.cat( + [chunks_by_flat_index[index] for index in prompt_order], dim=0) + + def prepare_inputs_for_generation( + self, + past_key_values: list[Sequence[torch.Tensor]], + inputs_embeds: torch.Tensor | None = None, + context: StepContext | None = None, + ): + named_states = context.named_state_caches + required_states = ( + GLM5_KDA_CONV_STATE, + GLM5_KDA_RECURRENT_STATE, + GLM5_KPOOL_TAIL_K_STATE, + GLM5_KPOOL_TAIL_SCORE_STATE, + ) + if named_states is None: + raise RuntimeError('GLM-5.3 requires named state caches.') + missing = [name for name in required_states if name not in named_states] + if missing: + raise RuntimeError( + f'GLM-5.3 is missing named state caches: {missing}.') + + # These specs are deliberately unlayered: their leading dimension is + # the compact KDA/full-attention row. Runtime storage leads with the + # request-state slot, so transpose once into layer-major views. + kda_conv = named_states[GLM5_KDA_CONV_STATE].transpose(0, 1) + kda_recurrent = named_states[ + GLM5_KDA_RECURRENT_STATE].transpose(0, 1) + tail_k = named_states[GLM5_KPOOL_TAIL_K_STATE].transpose(0, 1) + tail_score = named_states[ + GLM5_KPOOL_TAIL_SCORE_STATE].transpose(0, 1) + linear_caches = list(zip(kda_conv, kda_recurrent)) + kpool_tail_states = list(zip(tail_k, tail_score)) + full_caches = list(past_key_values) + interleaved_caches = [] + for layer_idx in range(self.config.num_hidden_layers): + if is_glm5_kda_layer(self.config, layer_idx): + interleaved_caches.append(linear_caches.pop(0)) + else: + interleaved_caches.append(full_caches.pop(0)) + if linear_caches or full_caches: + raise RuntimeError( + 'GLM-5.3 cache counts do not match its hybrid layer map.') + + input_ids = context.input_ids + pixel_values = None + grid_thw = None + vision_groups = None + vision_prompt_order = None + multimodal_mask = None + if context.input_multimodals is not None: + (vision_groups, vision_prompt_order, + mm_inputs) = self._group_vision_inputs( + context.input_multimodals) + if vision_groups: + multimodal_mask = self._multimodal_token_mask( + input_ids, mm_inputs) + + vision_embeddings = context.input_embeddings + vision_embedding_indexing = context.input_embedding_indexing + if vision_embeddings is not None and len(vision_embeddings) > 0: + if inputs_embeds is None: + inputs_embeds = self.get_input_embeddings()(input_ids) + inputs_embeds[:, + vision_embedding_indexing, :] = vision_embeddings.to( + inputs_embeds) + + return dict(input_ids=input_ids, + position_ids=context.position_ids, + past_key_values=interleaved_caches, + attn_metadata=context.attn_metadata, + inputs_embeds=inputs_embeds, + state_ids=context.state_offsets, + kpool_tail_states=kpool_tail_states, + pixel_values=pixel_values, + grid_thw=grid_thw, + vision_groups=vision_groups, + vision_prompt_order=vision_prompt_order, + multimodal_mask=multimodal_mask) + + @staticmethod + def _layer_idx(name: str) -> int | None: + match = re.search(r'\.layers\.(\d+)\.', name) + return None if match is None else int(match.group(1)) + + def _load_weight_attention(self, name: str, loaded_weight: torch.Tensor, + params_dict: dict[str, nn.Parameter], + update_pe_mapping: list): + """Retain KV-B while deriving the absorbed KC/VC views from it.""" + if name.endswith('.kv_b_proj.weight'): + load_weight(params_dict[name], loaded_weight) + return super()._load_weight_attention(name, loaded_weight, params_dict, + update_pe_mapping) + + def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]): + """Load both towers through LMDeploy's TP-aware weight loaders.""" + stacked_params_mapping = [ + ('.gate_up_proj', '.gate_proj', 0), + ('.gate_up_proj', '.up_proj', 1), + ] + kda_params_mapping = [ + ('.qkv_proj', '.q_proj', 'q'), + ('.qkv_proj', '.k_proj', 'k'), + ('.qkv_proj', '.v_proj', 'v'), + ('.qkv_conv1d', '.q_conv1d', 'q'), + ('.qkv_conv1d', '.k_conv1d', 'k'), + ('.qkv_conv1d', '.v_conv1d', 'v'), + ] + expert_params_mapping = [] + for expert_id in range(self.config.n_routed_experts): + expert_params_mapping.extend([ + ('.experts.gate_up', f'.experts.{expert_id}.gate_proj', + expert_id, 'gate'), + ('.experts.gate_up', f'.experts.{expert_id}.up_proj', + expert_id, 'up'), + ('.experts.down', f'.experts.{expert_id}.down_proj', expert_id, + 'down'), + ]) + + params_dict = dict(self.named_parameters()) + for checkpoint_name, loaded_weight in weights: + is_visual_weight = checkpoint_name.startswith('model.visual.') + if checkpoint_name.startswith('model.visual.'): + if getattr(self.visual, '_is_dummy_mod', False): + continue + name = checkpoint_name.replace('model.visual.', 'visual.', 1) + else: + name = checkpoint_name.replace('model.language_model.', + 'model.', 1) + layer_idx = self._layer_idx(name) + if layer_idx is not None and layer_idx >= self.config.num_hidden_layers: + # MTP layer 45 is a separate milestone. + continue + if 'rotary_emb.' in name: + continue + if self.config.tie_word_embeddings and name == 'lm_head.weight': + continue + + if is_visual_weight and '.attn.qkv.' in name: + param = params_dict[name] + query, key, value = param.weight_spliter(loaded_weight) + load_weight(param, query, shard_id='q') + load_weight(param, key, shard_id='k') + load_weight(param, value, shard_id='v') + continue + + if '.experts.' in name: + self._load_weight_experts(name, loaded_weight, params_dict, + expert_params_mapping) + continue + + is_kda = (layer_idx is not None + and is_glm5_kda_layer(self.config, layer_idx) + and '.self_attn.' in name) + if is_kda: + for param_name, weight_name, shard_id in kda_params_mapping: + if weight_name not in name: + continue + name = name.replace(weight_name, param_name) + load_weight(params_dict[name], + loaded_weight, + shard_id=shard_id) + break + else: + load_weight(params_dict[name], loaded_weight) + continue + + if layer_idx is not None and '.self_attn.' in name: + if '.self_attn.indexer.' in name: + # KPool projections are replicated (not TP-sharded), and + # their seven checkpoint tensors map one-to-one. + load_weight(params_dict[name], loaded_weight) + continue + self._load_weight_attention(name, + loaded_weight, + params_dict, + update_pe_mapping=[]) + continue + + for param_name, weight_name, shard_id in stacked_params_mapping: + if weight_name not in name: + continue + name = name.replace(weight_name, param_name) + load_weight(params_dict[name], + loaded_weight, + shard_id=shard_id) + break + else: + load_weight(params_dict[name], loaded_weight) diff --git a/lmdeploy/pytorch/models/module_map.py b/lmdeploy/pytorch/models/module_map.py index 6fefd3759f..15b7298574 100644 --- a/lmdeploy/pytorch/models/module_map.py +++ b/lmdeploy/pytorch/models/module_map.py @@ -64,6 +64,12 @@ # glm5 MODULE_MAP.update({'GlmMoeDsaForCausalLM': f'{LMDEPLOY_PYTORCH_MODEL_PATH}.glm_moe_dsa.GlmMoeDsaForCausalLM'}) +# glm5.3 flash +MODULE_MAP.update({ + 'Glm5NextForConditionalGeneration': + f'{LMDEPLOY_PYTORCH_MODEL_PATH}.glm5_next.Glm5NextForConditionalGeneration' +}) + # glm5 mtp MODULE_MAP.update({'GlmMoeDsaMTPModel': f'{LMDEPLOY_PYTORCH_MODEL_PATH}.glm_moe_dsa_mtp.GlmMoeDsaMTPModel'}) diff --git a/lmdeploy/pytorch/nn/__init__.py b/lmdeploy/pytorch/nn/__init__.py index c167e588ec..f321e96cdd 100644 --- a/lmdeploy/pytorch/nn/__init__.py +++ b/lmdeploy/pytorch/nn/__init__.py @@ -5,11 +5,14 @@ from .attention import Attention, FlashAttention # noqa: F401 from .embedding import ParallelEmbedding, ParallelLMHead # noqa: F401 from .hc_prepost import HcPrePost # noqa: F401 -from .norm import LayerNorm, RMSNorm, rms_scale # noqa: F401 +from .kda import Kda # noqa: F401 +from .kpool import KPoolIndexer # noqa: F401 +from .norm import FP32LayerNorm, LayerNorm, RMSNorm, rms_scale # noqa: F401 from .rotary_embedding import ( ApplyRotaryEmb, # noqa: F401 RopeType, # noqa: F401 YarnParameters, # noqa: F401 + apply_rotary_pos_emb_fp32, # noqa: F401 build_rotary_embedding, # noqa: F401 build_rotary_embedding_from_config, # noqa: F401 build_rotary_params, # noqa: F401 diff --git a/lmdeploy/pytorch/nn/attention.py b/lmdeploy/pytorch/nn/attention.py index 39db348c9e..d8b52f0a25 100644 --- a/lmdeploy/pytorch/nn/attention.py +++ b/lmdeploy/pytorch/nn/attention.py @@ -90,6 +90,36 @@ def _lazy_init(self, device): self.impl.set_alibi_slopes(alibi_slopes) self.alibi_ready = True + def fill_and_flatten_latent_kv_cache( + self, + key: torch.Tensor, + k_cache: torch.Tensor, + attn_metadata: AttentionMetadata, + out_dtype: torch.dtype = None, + k_scales_zeros: torch.Tensor = None, + v_scales_zeros: torch.Tensor = None, + ) -> torch.Tensor: + """Append latent KV and return the complete request-major prefill KV.""" + self._lazy_init(key.device) + + quant_policy = attn_metadata.quant_policy + if quant_policy in (QuantPolicy.FP8, QuantPolicy.FP8_E5M2): + if self.k_scale.device != key.device: + self.k_scale = self.k_scale.to(device=key.device, non_blocking=True) + if self.v_scale.device != key.device: + self.v_scale = self.v_scale.to(device=key.device, non_blocking=True) + k_scales_zeros = self.k_scale + v_scales_zeros = self.v_scale + + return self.impl.fill_and_flatten_latent_kv_cache( + key, + k_cache, + attn_metadata, + out_dtype=out_dtype, + k_scales_zeros=k_scales_zeros, + v_scales_zeros=v_scales_zeros, + ) + def forward( self, query: torch.Tensor, diff --git a/lmdeploy/pytorch/nn/kda.py b/lmdeploy/pytorch/nn/kda.py new file mode 100644 index 0000000000..4dbd01599a --- /dev/null +++ b/lmdeploy/pytorch/nn/kda.py @@ -0,0 +1,48 @@ +# Copyright (c) OpenMMLab. All rights reserved. +from typing import Any + +import torch +from torch import nn + +from lmdeploy.pytorch.backends import OpType, get_backend + + +class Kda(nn.Module): + """Backend-dispatched Kimi Delta Attention recurrence.""" + + def __init__(self): + super().__init__() + builder = get_backend().get_layer_impl_builder(OpType.Kda) + self.impl = builder.build() + + def forward( + self, + mixed_qkv: torch.Tensor, + raw_gate: torch.Tensor, + raw_beta: torch.Tensor, + conv_weight: torch.Tensor, + conv_bias: torch.Tensor | None, + a_log: torch.Tensor, + dt_bias: torch.Tensor, + conv_state: torch.Tensor, + recurrent_state: torch.Tensor, + metadata: Any, + num_heads: int, + head_dim: int, + lower_bound: float, + ) -> torch.Tensor: + return self.impl.forward( + mixed_qkv=mixed_qkv, + raw_gate=raw_gate, + raw_beta=raw_beta, + conv_weight=conv_weight, + conv_bias=conv_bias, + a_log=a_log, + dt_bias=dt_bias, + conv_state=conv_state, + recurrent_state=recurrent_state, + metadata=metadata, + num_heads=num_heads, + head_dim=head_dim, + lower_bound=lower_bound, + ) diff --git a/lmdeploy/pytorch/nn/kpool.py b/lmdeploy/pytorch/nn/kpool.py new file mode 100644 index 0000000000..1bae749ac8 --- /dev/null +++ b/lmdeploy/pytorch/nn/kpool.py @@ -0,0 +1,897 @@ +# Copyright (c) OpenMMLab. All rights reserved. +"""Reusable Torch semantics for the DSA KPool indexer. + +KPool has two different kinds of runtime data. Closed pools are pageable and +belong in the named DSA index cache. The unfinished per-request tail is +sequence state and must be supplied by the caller; this module deliberately +does not hide it in mutable module tensors. + +The functions here are a device-agnostic correctness path. CUDA backends can +replace compression, FP8 scoring, and top-k with fused kernels while retaining +these input/output contracts. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Literal + +import torch +import torch.nn.functional as F +from torch import Tensor, nn + +from lmdeploy.pytorch.consts import DSA_INDEX_SCALE_BYTES +from lmdeploy.pytorch.nn.linear import build_colwise_linear +from lmdeploy.pytorch.nn.norm import FP32LayerNorm + +KPOOL_PAGE_SIZE = 64 +KPOOL_FP8_MAX = 448.0 +KPOOL_NORM_EPS = 1e-6 +KPOOL_SCALE_FORMAT = 'ue8m0' +KPOOL_INDEXER_PARAMETER_NAMES = ( + 'index_kpool_compress_ape', + 'index_kpool_compress_gate', + 'k_norm.weight', + 'k_norm.bias', + 'weights_proj.weight', + 'wk.weight', + 'wq_b.weight', +) + + +def _validate_pool_geometry(pool_size: int, topk: int | None = None) -> None: + if pool_size <= 1: + raise ValueError(f'KPool pool_size must be greater than one, got {pool_size}.') + if KPOOL_PAGE_SIZE % pool_size: + raise ValueError(f'KPool pool_size must divide page size {KPOOL_PAGE_SIZE}, got {pool_size}.') + if topk is not None and (topk <= 0 or topk % pool_size): + raise ValueError(f'KPool topk must be positive and divisible by pool_size, got topk={topk}, ' + f'pool_size={pool_size}.') + + +@dataclass(frozen=True) +class KPoolUpdate: + """Closed pools and the unfinished tail produced by one request chunk. + + ``closed_group_ids`` are request-local logical pool ids. ``tail_keys`` and + ``tail_scores`` start at ``tail_logical_start`` and must be persisted as + request state before the next chunk or decode token. + """ + + closed_group_ids: Tensor + closed_keys: Tensor + closed_scores: Tensor + tail_keys: Tensor + tail_scores: Tensor + tail_logical_start: int + + +@dataclass(frozen=True) +class KPoolDecodeUpdate: + """Fixed-shape KPool update used by batched autoregressive decode. + + CUDA Graph decode pads requests to a capture bucket, so this contract + keeps every tensor batch-shaped and represents inactive rows with + ``valid_state=False`` instead of materializing dynamic index tensors. + """ + + closed_keys: Tensor + closed_scores: Tensor + group_ids: Tensor + should_close: Tensor + next_tail_keys: Tensor + next_tail_scores: Tensor + safe_state_ids: Tensor + valid_state: Tensor + + +class KPoolIndexer(nn.Module): + """Replicated KPool parameter layer shared by all attention-TP ranks. + + The seven parameter names and dtypes match the GLM-5.3/SGLang checkpoint + contract. Query/key rotary handling and cache ownership stay with the + model/backend because they depend on model geometry and request metadata. + """ + + def __init__( + self, + hidden_size: int, + index_n_heads: int, + index_head_dim: int, + index_topk: int, + q_lora_rank: int, + index_kpool: int, + norm_eps: float = KPOOL_NORM_EPS, + scale_fmt: str | None = KPOOL_SCALE_FORMAT, + dtype: torch.dtype | None = None, + device: torch.device | str | None = None, + prefix: str = '', + key_norm: nn.Module | None = None, + ) -> None: + super().__init__() + _validate_pool_geometry(index_kpool, index_topk) + if index_head_dim <= 0 or index_head_dim & (index_head_dim - 1): + raise ValueError(f'KPool index_head_dim must be a positive power of two, got {index_head_dim}.') + if dtype is None: + dtype = torch.bfloat16 + + self.hidden_size = hidden_size + self.n_heads = index_n_heads + self.head_dim = index_head_dim + self.index_topk = index_topk + self.q_lora_rank = q_lora_rank + self.index_kpool = index_kpool + self.softmax_scale = index_head_dim**-0.5 + self.scale_fmt = scale_fmt + + def add_prefix(name: str) -> str: + return f'{prefix}.{name}' if prefix else name + + # SGLang's ReplicatedLinear contract: none of these projections is TP + # sharded, even when the surrounding MLA attention runs with TP=8. + self.wq_b = build_colwise_linear( + q_lora_rank, + index_n_heads * index_head_dim, + bias=False, + dtype=dtype, + device=device, + is_tp=False, + quant_config=None, + check_dist=False, + prefix=add_prefix('wq_b'), + ) + self.wk = build_colwise_linear( + hidden_size, + index_head_dim, + bias=False, + dtype=dtype, + device=device, + is_tp=False, + quant_config=None, + check_dist=False, + prefix=add_prefix('wk'), + ) + # The checkpoint stores BF16, but SGLang promotes these parameters to + # FP32 before inference. LMDeploy's loader performs the same cast. + self.weights_proj = build_colwise_linear( + hidden_size, + index_n_heads, + bias=False, + dtype=torch.float32, + device=device, + is_tp=False, + quant_config=None, + check_dist=False, + prefix=add_prefix('weights_proj'), + ) + # The device-agnostic KPool owner keeps a Torch reference default. + # Models that require a platform-exact provider can inject the same + # parameter-shaped component without changing checkpoint names. + self.k_norm = (key_norm if key_norm is not None else FP32LayerNorm( + index_head_dim, norm_eps, device=device)) + self.index_kpool_compress_ape = nn.Parameter( + torch.zeros(index_kpool, index_head_dim, dtype=torch.float32, device=device), + requires_grad=False, + ) + self.index_kpool_compress_gate = nn.Parameter( + torch.empty(index_head_dim, hidden_size, dtype=dtype, device=device), + requires_grad=False, + ) + + def project_query(self, q_lora: Tensor) -> Tensor: + """Project the latent query to ``[..., index_n_heads, head_dim]``.""" + query = self.wq_b(q_lora) + return query.unflatten(-1, (self.n_heads, self.head_dim)) + + def project_key(self, hidden_states: Tensor) -> Tensor: + """Project and normalize one shared index key per token.""" + return self.k_norm(self.wk(hidden_states)) + + def project_compress_score(self, hidden_states: Tensor) -> Tensor: + """Return per-slot, per-dimension compression gates.""" + return F.linear(hidden_states, self.index_kpool_compress_gate) + + def project_head_gate(self, hidden_states: Tensor) -> Tensor: + """Return FP32 head gates including SGLang's ``num_heads**-0.5``.""" + return self.weights_proj(hidden_states.float()) * self.n_heads**-0.5 + + def quantize_fp8(self, values: Tensor) -> tuple[Tensor, Tensor]: + """Quantize index vectors with the checkpoint's UE8M0 scale rule.""" + return kpool_quantize_fp8( + values, + block_size=self.head_dim, + round_scale=self.scale_fmt is not None, + ) + + +def kpool_normalized_hadamard(values: Tensor) -> Tensor: + """Apply the normalized Walsh-Hadamard transform on the last dimension.""" + width = values.size(-1) + if width <= 0 or width & (width - 1): + raise ValueError(f'Hadamard width must be a positive power of two, got {width}.') + + output = values.float() + stride = 1 + while stride < width: + shape = output.shape[:-1] + (-1, 2, stride) + paired = output.reshape(shape) + left, right = paired.unbind(dim=-2) + output = torch.cat((left + right, left - right), dim=-1).reshape_as(output) + stride *= 2 + return output * width**-0.5 + + +def kpool_rotate_query(query: Tensor) -> Tensor: + """Match the BF16-preserving query rotation used before FP8 quantization.""" + return kpool_normalized_hadamard(query).to(query.dtype) + + +def kpool_partition_update( + chunk_keys: Tensor, + chunk_scores: Tensor, + history_length: int, + pool_size: int, + tail_keys: Tensor | None = None, + tail_scores: Tensor | None = None, +) -> KPoolUpdate: + """Assemble arbitrary-length input with a prior tail into closed pools. + + This is the state-free equivalent of SGLang's extend/decode tail ring. The + caller owns persistence of the returned tail and supplies it on the next + invocation. + """ + _validate_pool_geometry(pool_size) + if history_length < 0: + raise ValueError(f'history_length must be non-negative, got {history_length}.') + if chunk_keys.ndim != 2 or chunk_scores.shape != chunk_keys.shape: + raise ValueError('chunk_keys and chunk_scores must have the same [tokens, head_dim] shape.') + + previous_tail_len = history_length % pool_size + if tail_keys is None: + tail_keys = chunk_keys[:0] + if tail_scores is None: + tail_scores = chunk_scores[:0] + expected_tail_shape = (previous_tail_len, chunk_keys.size(1)) + if tail_keys.shape != expected_tail_shape or tail_scores.shape != expected_tail_shape: + raise ValueError('The supplied tail must contain exactly history_length % pool_size rows; ' + f'expected {expected_tail_shape}, got keys={tuple(tail_keys.shape)}, ' + f'scores={tuple(tail_scores.shape)}.') + if tail_keys.device != chunk_keys.device or tail_scores.device != chunk_scores.device: + raise ValueError('Chunk and tail tensors must be on the same device.') + + all_keys = torch.cat((tail_keys, chunk_keys), dim=0) + all_scores = torch.cat((tail_scores, chunk_scores), dim=0) + num_closed = all_keys.size(0) // pool_size + num_closed_tokens = num_closed * pool_size + first_group_id = (history_length - previous_tail_len) // pool_size + closed_group_ids = torch.arange( + first_group_id, + first_group_id + num_closed, + dtype=torch.int64, + device=chunk_keys.device, + ) + closed_shape = (num_closed, pool_size, chunk_keys.size(1)) + closed_keys = all_keys[:num_closed_tokens].reshape(closed_shape) + closed_scores = all_scores[:num_closed_tokens].reshape(closed_shape) + new_tail_keys = all_keys[num_closed_tokens:] + new_tail_scores = all_scores[num_closed_tokens:] + tail_logical_start = history_length + chunk_keys.size(0) - new_tail_keys.size(0) + return KPoolUpdate( + closed_group_ids=closed_group_ids, + closed_keys=closed_keys, + closed_scores=closed_scores, + tail_keys=new_tail_keys, + tail_scores=new_tail_scores, + tail_logical_start=tail_logical_start, + ) + + +def kpool_decode_update( + keys: Tensor, + scores: Tensor, + tail_key_state: Tensor, + tail_score_state: Tensor, + state_ids: Tensor, + history_lengths: Tensor, + pool_size: int, +) -> KPoolDecodeUpdate: + """Build a fixed-shape, graph-safe update for one-token decode. + + State slot zero is LMDeploy's reserved dummy slot. Invalid CUDA Graph + padding rows therefore read and write slot zero without affecting a live + request. Every returned tensor has a shape determined only by the graph + capture bucket. + """ + _validate_pool_geometry(pool_size) + if keys.ndim != 2 or scores.shape != keys.shape: + raise ValueError( + 'keys and scores must have the same [batch, head_dim] shape.') + if tail_key_state.shape != tail_score_state.shape: + raise ValueError('KPool key and score tail states must match.') + if tail_key_state.ndim != 3: + raise ValueError( + 'KPool tail state must have shape [state, pool_size, head_dim].') + if tail_key_state.shape[1:] != (pool_size, keys.size(1)): + raise ValueError( + 'KPool tail state geometry does not match decode inputs.') + batch_size = keys.size(0) + if state_ids.shape != (batch_size, ) or history_lengths.shape != ( + batch_size, ): + raise ValueError( + 'state_ids and history_lengths must contain one value per row.') + + safe_state_ids = state_ids.to(torch.int64).clamp_min(0) + valid_state = state_ids >= 0 + history_lengths = history_lengths.to(torch.int64).clamp_min(0) + previous_tail_lengths = torch.remainder(history_lengths, pool_size) + previous_keys = tail_key_state.index_select(0, safe_state_ids) + previous_scores = tail_score_state.index_select(0, safe_state_ids) + + slots = torch.arange( + pool_size, dtype=torch.int64, device=keys.device)[None, :, None] + previous_valid = slots < previous_tail_lengths[:, None, None] + closed_keys = torch.where(previous_valid, previous_keys, + torch.zeros_like(previous_keys)) + closed_scores = torch.where(previous_valid, previous_scores, + torch.zeros_like(previous_scores)) + insert_at = previous_tail_lengths[:, None, None].expand( + -1, 1, keys.size(1)) + closed_keys.scatter_(1, insert_at, keys[:, None, :]) + closed_scores.scatter_(1, insert_at, scores[:, None, :]) + + should_close = valid_state & (previous_tail_lengths == pool_size - 1) + next_tail_keys = torch.where(should_close[:, None, None], + torch.zeros_like(closed_keys), closed_keys) + next_tail_scores = torch.where(should_close[:, None, None], + torch.zeros_like(closed_scores), + closed_scores) + # Invalid graph-padding rows preserve the reserved dummy state. + next_tail_keys = torch.where(valid_state[:, None, None], next_tail_keys, + previous_keys) + next_tail_scores = torch.where(valid_state[:, None, None], + next_tail_scores, previous_scores) + group_ids = torch.div( + history_lengths, pool_size, rounding_mode='floor') + return KPoolDecodeUpdate( + closed_keys=closed_keys, + closed_scores=closed_scores, + group_ids=group_ids, + should_close=should_close, + next_tail_keys=next_tail_keys, + next_tail_scores=next_tail_scores, + safe_state_ids=safe_state_ids, + valid_state=valid_state, + ) + + +KPoolCompressMode = Literal['extend', 'decode'] + + +def _validate_compress_inputs(closed_keys: Tensor, closed_scores: Tensor, + ape: Tensor) -> None: + if closed_keys.ndim != 3 or closed_scores.shape != closed_keys.shape: + raise ValueError( + 'closed_keys and closed_scores must have the same ' + '[groups, pool_size, head_dim] shape.') + if ape.shape != closed_keys.shape[1:]: + raise ValueError( + f'ape must have shape {tuple(closed_keys.shape[1:])}, ' + f'got {tuple(ape.shape)}.') + + +def _finish_compressed_pool(pooled: Tensor) -> Tensor: + # SGLang rounds both the pooled vector and the rotated vector through BF16. + pooled = pooled.to(torch.bfloat16).float() + return kpool_normalized_hadamard(pooled).to(torch.bfloat16) + + +def kpool_compress_online(closed_keys: Tensor, closed_scores: Tensor, + ape: Tensor) -> Tensor: + """Use the online recurrence from SGLang's extend assembly kernel. + + The explicit slot loop matches SGLang's compression kernel reduction + order. A vectorized max/exp/sum is mathematically equivalent but can move + values across the FP8 quantization boundary after different FP32 rounds. + """ + _validate_compress_inputs(closed_keys, closed_scores, ape) + + groups, pool_size, head_dim = closed_keys.shape + max_score = torch.full( + (groups, head_dim), + -float('inf'), + dtype=torch.float32, + device=closed_keys.device, + ) + denominator = torch.zeros_like(max_score) + accumulator = torch.zeros_like(max_score) + for slot in range(pool_size): + score = closed_scores[:, slot].float() + ape[slot].float() + new_max = torch.maximum(max_score, score) + rescale = torch.exp(max_score - new_max) + probability = torch.exp(score - new_max) + key = closed_keys[:, slot].float() + denominator = denominator * rescale + probability + accumulator = accumulator * rescale + key * probability + max_score = new_max + return _finish_compressed_pool(accumulator / denominator) + + +def kpool_compress_two_pass(closed_keys: Tensor, closed_scores: Tensor, + ape: Tensor) -> Tensor: + """Use the max-then-sum order from SGLang's decode close-pool kernel.""" + _validate_compress_inputs(closed_keys, closed_scores, ape) + groups, pool_size, head_dim = closed_keys.shape + max_score = torch.full( + (groups, head_dim), + -float('inf'), + dtype=torch.float32, + device=closed_keys.device, + ) + for slot in range(pool_size): + score = closed_scores[:, slot].float() + ape[slot].float() + max_score = torch.maximum(max_score, score) + + denominator = torch.zeros_like(max_score) + accumulator = torch.zeros_like(max_score) + for slot in range(pool_size): + score = closed_scores[:, slot].float() + ape[slot].float() + probability = torch.exp(score - max_score) + key = closed_keys[:, slot].float() + denominator = denominator + probability + accumulator = accumulator + key * probability + return _finish_compressed_pool(accumulator / denominator) + + +def kpool_compress( + closed_keys: Tensor, + closed_scores: Tensor, + ape: Tensor, + *, + mode: KPoolCompressMode, +) -> Tensor: + """Dispatch the explicit SGLang compression order for a forward mode.""" + if mode == 'extend': + return kpool_compress_online(closed_keys, closed_scores, ape) + if mode == 'decode': + return kpool_compress_two_pass(closed_keys, closed_scores, ape) + raise ValueError( + f'KPool compression mode must be extend or decode, got {mode!r}.') + + +def kpool_quantize_fp8( + values: Tensor, + block_size: int = 128, + round_scale: bool = True, +) -> tuple[Tensor, Tensor]: + """Block-quantize to E4M3FN using SGLang's scale and clamp semantics.""" + if block_size <= 0 or values.size(-1) % block_size: + raise ValueError(f'Last dimension must be divisible by block_size, got shape={tuple(values.shape)}, ' + f'block_size={block_size}.') + grouped = values.float().unflatten(-1, (-1, block_size)) + absmax = grouped.abs().amax(dim=-1).clamp_min(1e-4) + scale = absmax / KPOOL_FP8_MAX + if round_scale: + scale = torch.exp2(torch.ceil(torch.log2(scale))) + quantized = (grouped / scale.unsqueeze(-1)).clamp(-KPOOL_FP8_MAX, KPOOL_FP8_MAX) + quantized = quantized.flatten(-2).to(torch.float8_e4m3fn) + return quantized, scale.to(torch.float32) + + +def kpool_score( + query_fp8: Tensor, + query_scale: Tensor, + pooled_key_fp8: Tensor, + pooled_key_scale: Tensor, + head_gate: Tensor, + softmax_scale: float | None = None, +) -> Tensor: + """Reference FP8 MQA logits for pooled history. + + ``head_gate`` is the output of :meth:`KPoolIndexer.project_head_gate`, so it + already includes the ``num_heads**-0.5`` factor. Query and pooled-key + scales are applied exactly once here. + """ + if query_fp8.ndim != 3: + raise ValueError('query_fp8 must have shape [rows, heads, head_dim].') + if pooled_key_fp8.ndim != 2 or pooled_key_fp8.size(1) != query_fp8.size(2): + raise ValueError('pooled_key_fp8 must have shape [groups, query_head_dim].') + rows, heads, head_dim = query_fp8.shape + if head_gate.shape != (rows, heads): + raise ValueError(f'head_gate must have shape {(rows, heads)}, got {tuple(head_gate.shape)}.') + if query_scale.shape == (rows, heads, 1): + query_scale = query_scale.squeeze(-1) + if query_scale.shape != (rows, heads): + raise ValueError(f'query_scale must have shape {(rows, heads)} or {(rows, heads, 1)}, ' + f'got {tuple(query_scale.shape)}.') + if pooled_key_scale.shape == (pooled_key_fp8.size(0), 1): + pooled_key_scale = pooled_key_scale.squeeze(-1) + if pooled_key_scale.shape != (pooled_key_fp8.size(0), ): + raise ValueError('pooled_key_scale must contain one scale per pooled key.') + if softmax_scale is None: + softmax_scale = head_dim**-0.5 + + query_weight = head_gate.float() * query_scale.float() * softmax_scale + per_head_logits = torch.einsum( + 'rhd,kd->rhk', query_fp8.float(), pooled_key_fp8.float()) + per_head_logits = per_head_logits.clamp_min_(0) + logits = torch.einsum('rhk,rh->rk', per_head_logits, query_weight) + return logits * pooled_key_scale.float().unsqueeze(0) + + +def kpool_pooled_block_offsets(token_block_offsets: Tensor, pool_size: int) -> Tensor: + """Build the pooled page table by selecting every ``pool_size`` token page.""" + _validate_pool_geometry(pool_size) + if token_block_offsets.ndim < 1: + raise ValueError('token_block_offsets must have at least one dimension.') + columns = torch.arange(0, token_block_offsets.size(-1), pool_size, device=token_block_offsets.device) + return token_block_offsets.index_select(-1, columns) + + +def kpool_pooled_write_locations( + token_block_offsets: Tensor, + group_ids: Tensor, + pool_size: int, + page_size: int = KPOOL_PAGE_SIZE, +) -> Tensor: + """Map request-local logical pool ids to packed physical cache slots.""" + _validate_pool_geometry(pool_size) + if page_size != KPOOL_PAGE_SIZE: + raise ValueError(f'KPool currently requires page_size={KPOOL_PAGE_SIZE}, got {page_size}.') + if token_block_offsets.ndim != 1 or group_ids.ndim != 1: + raise ValueError('token_block_offsets and group_ids must both be one-dimensional.') + group_ids = group_ids.to(torch.int64) + page_group = torch.div(group_ids, page_size, rounding_mode='floor') + token_page_column = page_group * pool_size + if token_page_column.numel() and int(token_page_column.max()) >= token_block_offsets.numel(): + raise ValueError('token_block_offsets is too short for the requested logical pool ids.') + physical_page = token_block_offsets.index_select(0, token_page_column) + return physical_page.to(torch.int64) * page_size + torch.remainder(group_ids, page_size) + + +def kpool_packed_cache_views( + packed_cache: Tensor, + head_dim: int, +) -> tuple[Tensor, Tensor]: + """Expose FP8 values and FP32 scales from one packed DSA cache row. + + The byte layout intentionally matches the existing DeepGEMM DSA cache: + every page stores all 64 value rows first, followed by all 64 scales. + """ + if packed_cache.dtype != torch.uint8: + raise TypeError( + 'Packed KPool cache must be uint8, ' + f'got {packed_cache.dtype}.') + if packed_cache.dim() != 4 or packed_cache.size(2) != 1: + raise ValueError( + 'Packed KPool cache must have shape ' + '[num_blocks, entries, 1, head_dim + 4].') + packed_width = head_dim + DSA_INDEX_SCALE_BYTES + if packed_cache.size(-1) != packed_width: + raise ValueError( + f'Packed KPool cache last dim must be {packed_width}, ' + f'got {packed_cache.size(-1)}.') + + num_blocks, entries_per_block = packed_cache.shape[:2] + flat = packed_cache.view(num_blocks, -1) + value_bytes = entries_per_block * head_dim + scale_bytes = entries_per_block * DSA_INDEX_SCALE_BYTES + values = flat[:, :value_bytes].view(torch.float8_e4m3fn).view( + num_blocks, entries_per_block, head_dim) + scales = flat[:, value_bytes:value_bytes + scale_bytes].view( + torch.float32).view(num_blocks, entries_per_block, 1) + return values, scales + + +def kpool_write_packed_cache( + packed_cache: Tensor, + token_block_offsets: Tensor, + group_ids: Tensor, + pooled_key_fp8: Tensor, + pooled_key_scale: Tensor, + pool_size: int, +) -> None: + """Write closed request-local pools into their pageable packed cache.""" + if pooled_key_fp8.ndim != 2: + raise ValueError('pooled_key_fp8 must have shape [groups, head_dim].') + if pooled_key_scale.shape == (pooled_key_fp8.size(0), ): + pooled_key_scale = pooled_key_scale.unsqueeze(-1) + if pooled_key_scale.shape != (pooled_key_fp8.size(0), 1): + raise ValueError('pooled_key_scale must contain one scale per group.') + if group_ids.shape != (pooled_key_fp8.size(0), ): + raise ValueError('group_ids must contain one id per pooled key.') + values, scales = kpool_packed_cache_views( + packed_cache, pooled_key_fp8.size(-1)) + locations = kpool_pooled_write_locations( + token_block_offsets, + group_ids, + pool_size, + page_size=values.size(1), + ) + pages = torch.div(locations, values.size(1), rounding_mode='floor') + slots = torch.remainder(locations, values.size(1)) + values[pages, slots] = pooled_key_fp8.to(values.dtype) + scales[pages, slots] = pooled_key_scale.to(scales.dtype) + + +def kpool_write_packed_cache_batched( + packed_cache: Tensor, + token_block_offsets: Tensor, + group_ids: Tensor, + pooled_key_fp8: Tensor, + pooled_key_scale: Tensor, + pool_size: int, + valid: Tensor, +) -> None: + """Write at most one closed pool per fixed-shape decode row. + + Invalid rows target reserved cache block zero and keep its previous value. + This avoids dynamic boolean indexing and keeps CUDA Graph addresses and + launch geometry stable. + """ + if pooled_key_fp8.ndim != 2: + raise ValueError( + 'pooled_key_fp8 must have shape [batch, head_dim].') + batch_size = pooled_key_fp8.size(0) + if pooled_key_scale.shape == (batch_size, ): + pooled_key_scale = pooled_key_scale.unsqueeze(-1) + if pooled_key_scale.shape != (batch_size, 1): + raise ValueError('pooled_key_scale must contain one scale per row.') + if token_block_offsets.ndim != 2 or token_block_offsets.size(0) != batch_size: + raise ValueError( + 'token_block_offsets must have one page-table row per batch row.') + if group_ids.shape != (batch_size, ) or valid.shape != (batch_size, ): + raise ValueError('group_ids and valid must contain one value per row.') + + values, scales = kpool_packed_cache_views( + packed_cache, pooled_key_fp8.size(-1)) + group_ids = group_ids.to(torch.int64) + page_group = torch.div( + group_ids, values.size(1), rounding_mode='floor') + token_page_column = page_group * pool_size + token_page_column = token_page_column.clamp( + min=0, max=token_block_offsets.size(1) - 1) + physical_page = token_block_offsets.gather( + 1, token_page_column[:, None]).squeeze(1).to(torch.int64) + locations = physical_page * values.size(1) + torch.remainder( + group_ids, values.size(1)) + # Cache block zero is reserved by LMDeploy and is the shared dummy target. + locations = torch.where(valid, locations, torch.zeros_like(locations)) + pages = torch.div(locations, values.size(1), rounding_mode='floor') + slots = torch.remainder(locations, values.size(1)) + current_values = values[pages, slots] + current_scales = scales[pages, slots] + values[pages, slots] = torch.where( + valid[:, None], pooled_key_fp8.to(values.dtype), current_values) + scales[pages, slots] = torch.where( + valid[:, None], pooled_key_scale.to(scales.dtype), current_scales) + + +def kpool_read_packed_cache( + packed_cache: Tensor, + token_block_offsets: Tensor, + num_groups: int, + pool_size: int, +) -> tuple[Tensor, Tensor]: + """Gather a request's closed pools in logical order from paged storage.""" + if num_groups < 0: + raise ValueError(f'num_groups must be non-negative, got {num_groups}.') + head_dim = packed_cache.size(-1) - DSA_INDEX_SCALE_BYTES + values, scales = kpool_packed_cache_views(packed_cache, head_dim) + group_ids = torch.arange( + num_groups, dtype=torch.int64, device=packed_cache.device) + locations = kpool_pooled_write_locations( + token_block_offsets, + group_ids, + pool_size, + page_size=values.size(1), + ) + pages = torch.div(locations, values.size(1), rounding_mode='floor') + slots = torch.remainder(locations, values.size(1)) + return values[pages, slots], scales[pages, slots] + + +def kpool_selected_token_counts(seq_lens: Tensor, topk: int, pool_size: int) -> Tensor: + """Return selected history plus always-selected ragged tail token counts.""" + _validate_pool_geometry(pool_size, topk) + full_pool_tokens = torch.div(seq_lens, pool_size, rounding_mode='floor') * pool_size + return full_pool_tokens.clamp(max=topk) + seq_lens - full_pool_tokens + + +def _map_logical_indices( + logical: Tensor, + valid: Tensor, + page_table: Tensor | None, + topk_offsets: Tensor | None, +) -> Tensor: + if page_table is not None and topk_offsets is not None: + raise ValueError('page_table and topk_offsets are mutually exclusive.') + if page_table is not None: + if page_table.ndim != 2 or page_table.size(0) != logical.size(0): + raise ValueError('page_table must have one [logical_token] row per score row.') + if page_table.size(1) == 0: + if not valid.is_cuda and bool(valid.any()): + raise ValueError('A non-empty logical selection cannot use an empty page_table.') + return torch.full_like(logical, -1, dtype=torch.int32) + safe = logical.clamp(min=0, max=page_table.size(1) - 1) + output = page_table.gather(1, safe).to(torch.int32) + elif topk_offsets is not None: + if topk_offsets.ndim == 2 and topk_offsets.size(1) == 1: + topk_offsets = topk_offsets.squeeze(1) + if topk_offsets.shape != (logical.size(0), ): + raise ValueError('topk_offsets must contain one value per score row.') + output = (logical + topk_offsets.to(torch.int64).unsqueeze(1)).to(torch.int32) + else: + output = logical.to(torch.int32) + return torch.where(valid, output, torch.full_like(output, -1)) + + +def kpool_expand_selected_groups( + selected_groups: Tensor, + group_lengths: Tensor, + pool_size: int, + topk: int, + *, + seq_lens: Tensor | None = None, + page_table: Tensor | None = None, + topk_offsets: Tensor | None = None, + page_table_row_index: Tensor | None = None, + out_rows: int | None = None, +) -> Tensor: + """Expand preselected pool ids to tokens and append ragged tails. + + The returned width is ``topk`` without ``seq_lens`` and + ``topk + pool_size - 1`` when the always-selected tail is requested. + Entries are request-local logical token ids unless ``page_table`` or + ``topk_offsets`` requests the same transform used by SGLang. The order of + ``selected_groups`` is deliberately preserved because sparse-attention + reduction order is numerically observable. + """ + _validate_pool_geometry(pool_size, topk) + group_budget = topk // pool_size + if selected_groups.ndim != 2 or selected_groups.size(1) != group_budget: + raise ValueError( + 'selected_groups must have shape ' + f'[rows, {group_budget}], got {tuple(selected_groups.shape)}.') + rows = selected_groups.size(0) + if group_lengths.shape != (rows, ): + raise ValueError( + 'group_lengths must contain one value per selected-groups row.') + if out_rows is not None and out_rows < rows: + raise ValueError(f'out_rows must be at least {rows}, got {out_rows}.') + device = selected_groups.device + group_lengths = group_lengths.to(device=device, dtype=torch.int64) + if not group_lengths.is_cuda and bool((group_lengths < 0).any()): + raise ValueError('group_lengths must be non-negative.') + selected_groups = selected_groups.to(device=device, dtype=torch.int64) + + page_table_for_rows = page_table + if page_table_row_index is not None: + if page_table is None: + raise ValueError('page_table_row_index requires page_table.') + if page_table_row_index.shape != (rows, ): + raise ValueError('page_table_row_index must contain one index per score row.') + page_table_for_rows = page_table.index_select(0, page_table_row_index.to(page_table.device, torch.int64)) + + ranks = torch.arange(group_budget, dtype=torch.int64, device=device) + selected_valid = ranks.unsqueeze(0) < group_lengths.clamp( + max=group_budget).unsqueeze(1) + selected_valid &= selected_groups >= 0 + selected_valid &= selected_groups < group_lengths.unsqueeze(1) + + offsets = torch.arange(pool_size, dtype=torch.int64, device=device) + logical = (selected_groups.unsqueeze(-1) * pool_size + offsets).reshape(rows, topk) + expanded_valid = selected_valid.unsqueeze(-1).expand(-1, -1, pool_size).reshape(rows, topk) + expanded = _map_logical_indices(logical, expanded_valid, page_table_for_rows, topk_offsets) + + if seq_lens is None: + result = expanded + else: + if seq_lens.shape != (rows, ): + raise ValueError('seq_lens must contain one value per score row.') + seq_lens = seq_lens.to(device=device, dtype=torch.int64) + if (not seq_lens.is_cuda and bool((torch.div( + seq_lens, pool_size, rounding_mode='floor') != + group_lengths).any())): + raise ValueError('group_lengths must equal floor(seq_lens / pool_size) when appending KPool tails.') + + output_width = topk + pool_size - 1 + output_columns = torch.arange( + output_width, dtype=torch.int64, device=device).unsqueeze(0) + history_width = (group_lengths * pool_size).clamp(max=topk).unsqueeze(1) + history_valid = output_columns < history_width + safe_history_columns = output_columns.clamp(max=topk - 1).expand(rows, -1) + history_values = expanded.gather(1, safe_history_columns) + + tail_offset = output_columns - history_width + tail_count = torch.remainder(seq_lens, pool_size).unsqueeze(1) + tail_valid = (tail_offset >= 0) & (tail_offset < tail_count) + tail_logical = group_lengths.unsqueeze(1) * pool_size + tail_offset + tail_values = _map_logical_indices(tail_logical, tail_valid, page_table_for_rows, topk_offsets) + + result = torch.full( + (rows, output_width), -1, dtype=torch.int32, device=device) + result = torch.where(history_valid, history_values, result) + result = torch.where(tail_valid, tail_values, result) + + if out_rows is not None and out_rows != rows: + padded = torch.full((out_rows, result.size(1)), -1, dtype=result.dtype, device=result.device) + padded[:rows] = result + return padded + return result + + +def kpool_topk( + logits: Tensor, + group_lengths: Tensor, + pool_size: int, + topk: int, + *, + seq_lens: Tensor | None = None, + row_starts: Tensor | None = None, + page_table: Tensor | None = None, + topk_offsets: Tensor | None = None, + page_table_row_index: Tensor | None = None, + out_rows: int | None = None, +) -> Tensor: + """Torch reference for SGLang's pooled-group selection semantics. + + Rows no longer than the group budget bypass score selection and retain + chronological group order. Longer rows use an unsorted Torch top-k as a + set-level CPU reference. The CUDA production path supplies groups from + the shared byte-radix selector to :func:`kpool_expand_selected_groups`. + """ + _validate_pool_geometry(pool_size, topk) + if logits.ndim != 2 or group_lengths.shape != (logits.size(0), ): + raise ValueError( + 'logits must be [rows, groups] with one group_lengths value per row.') + rows, columns = logits.shape + if row_starts is None: + row_starts = torch.zeros( + rows, dtype=torch.int64, device=logits.device) + else: + if row_starts.shape != (rows, ): + raise ValueError('row_starts must contain one value per score row.') + row_starts = row_starts.to( + device=logits.device, dtype=torch.int64) + group_lengths = group_lengths.to(device=logits.device, dtype=torch.int64) + invalid_windows = ((group_lengths < 0) | (row_starts < 0) + | (row_starts + group_lengths > columns)) + if bool(invalid_windows.any()): + raise ValueError( + 'Each [row_start, row_start + group_length) window must be inside logits.') + + group_budget = topk // pool_size + chronological = torch.arange( + group_budget, dtype=torch.int64, device=logits.device).expand(rows, -1) + selected_groups = torch.where( + chronological < group_lengths.unsqueeze(1), + chronological, + torch.full_like(chronological, -1), + ) + + long_rows = torch.nonzero( + group_lengths > group_budget, as_tuple=False).flatten() + if long_rows.numel(): + column_ids = torch.arange( + columns, dtype=torch.int64, device=logits.device).unsqueeze(0) + valid_window = (column_ids >= row_starts.unsqueeze(1)) & ( + column_ids < (row_starts + group_lengths).unsqueeze(1)) + masked_logits = logits.float().masked_fill( + ~valid_window, float('-inf')) + long_logits = masked_logits.index_select(0, long_rows) + _, indices = torch.topk( + long_logits, group_budget, dim=-1, sorted=False) + selected_groups[long_rows] = indices - row_starts.index_select( + 0, long_rows).unsqueeze(1) + + return kpool_expand_selected_groups( + selected_groups, + group_lengths, + pool_size, + topk, + seq_lens=seq_lens, + page_table=page_table, + topk_offsets=topk_offsets, + page_table_row_index=page_table_row_index, + out_rows=out_rows, + ) diff --git a/lmdeploy/pytorch/nn/linear/default.py b/lmdeploy/pytorch/nn/linear/default.py index e17f50d76b..d62788329c 100644 --- a/lmdeploy/pytorch/nn/linear/default.py +++ b/lmdeploy/pytorch/nn/linear/default.py @@ -117,17 +117,20 @@ def update_weights(self): def _forward_default(self, x, all_reduce, tp_sizes): """Default forward implement.""" + bias = self.bias + if self.is_tp and not self.colwise and self.tp_rank != 0: + bias = None if self.tp_mode == TPMode.DP_TP: rank = self.tp_rank return self.impl.forward(x, self.weight, - self.bias, + bias, all_reduce, group=self.tp_group, rank=rank, scatter_size=tp_sizes) else: - return self.impl.forward(x, self.weight, self.bias, all_reduce, group=self.tp_group) + return self.impl.forward(x, self.weight, bias, all_reduce, group=self.tp_group) class MergedBaseLinear(BaseLinear): diff --git a/lmdeploy/pytorch/nn/moe/__init__.py b/lmdeploy/pytorch/nn/moe/__init__.py index 216b7c9f59..fe487a2462 100644 --- a/lmdeploy/pytorch/nn/moe/__init__.py +++ b/lmdeploy/pytorch/nn/moe/__init__.py @@ -24,6 +24,10 @@ def build_fused_moe( layer_idx: int = 0, act_func: Callable = None, prefix: str = '', + *, + fp32_acc: bool = False, + output_scale: float = 1.0, + use_deep_gemm: bool = False, ): """Fused moe builder.""" quant_method = None @@ -107,6 +111,9 @@ def build_fused_moe( all_reduce=all_reduce, layer_idx=layer_idx, act_func=act_func, + fp32_acc=fp32_acc, + output_scale=output_scale, + use_deep_gemm=use_deep_gemm, ) elif quant_method == 'compressed-tensors': if bias: diff --git a/lmdeploy/pytorch/nn/moe/blocked_fp8.py b/lmdeploy/pytorch/nn/moe/blocked_fp8.py index 74f5e26d9e..20339c211c 100644 --- a/lmdeploy/pytorch/nn/moe/blocked_fp8.py +++ b/lmdeploy/pytorch/nn/moe/blocked_fp8.py @@ -156,7 +156,10 @@ def __init__(self, device: torch.device | None = None, all_reduce: bool = True, layer_idx: int = 0, - act_func: Callable = None): + act_func: Callable = None, + fp32_acc: bool = False, + output_scale: float = 1.0, + use_deep_gemm: bool = False): device = device or torch.device('cpu') dtype = dtype or torch.float16 @@ -186,7 +189,10 @@ def __init__(self, fp8_dtype=fp8_dtype, num_max_dispatch_tokens_per_rank=deep_ep_max_tokens_per_rank, layer_idx=layer_idx, - custom_gateup_act=act_func is not None) + custom_gateup_act=act_func is not None, + fp32_acc=fp32_acc, + output_scale=output_scale, + use_deep_gemm=use_deep_gemm) self.impl.set_scale_fmt(scale_fmt) if self.ep_size > 1: diff --git a/lmdeploy/pytorch/nn/norm.py b/lmdeploy/pytorch/nn/norm.py index 745f8f96ad..f29e2d68c7 100644 --- a/lmdeploy/pytorch/nn/norm.py +++ b/lmdeploy/pytorch/nn/norm.py @@ -1,6 +1,7 @@ # Copyright (c) OpenMMLab. All rights reserved. import torch +import torch.nn.functional as F from torch import nn from lmdeploy.pytorch.distributed import get_dist_group, get_tp_world_rank @@ -25,6 +26,44 @@ def rms_scale(a: torch.Tensor, b: torch.Tensor, dim: int = -1, eps: float = 1e-6 return out.to(result_dtype) +class FP32LayerNorm(nn.Module): + """LayerNorm with FP32 parameters and accumulation. + + Some model components keep LayerNorm weights in FP32 even when the model + activation dtype is BF16. Keep that numerical contract in one reusable + module and cast only the returned activation back to its input dtype. + """ + + def __init__(self, + hidden_size: int, + eps: float = 1e-6, + bias: bool = True, + device: torch.device | str | None = None): + super().__init__() + self.hidden_size = hidden_size + self.eps = eps + self.weight = nn.Parameter(torch.ones(hidden_size, + dtype=torch.float32, + device=device), + requires_grad=False) + if bias: + self.bias = nn.Parameter(torch.zeros(hidden_size, + dtype=torch.float32, + device=device), + requires_grad=False) + else: + self.register_parameter('bias', None) + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + """Normalize in FP32 and preserve the activation dtype.""" + output = F.layer_norm(hidden_states.float(), + (self.hidden_size, ), + self.weight, + self.bias, + self.eps) + return output.to(hidden_states.dtype) + + class RMSNorm(nn.Module): """RMS Norm with add residual.""" diff --git a/lmdeploy/pytorch/nn/rotary_embedding.py b/lmdeploy/pytorch/nn/rotary_embedding.py index 150b37f8ba..b290c8c364 100644 --- a/lmdeploy/pytorch/nn/rotary_embedding.py +++ b/lmdeploy/pytorch/nn/rotary_embedding.py @@ -238,6 +238,30 @@ def build_rotary_embedding_from_config(config: PretrainedConfig, device: torch.d return build_rotary_embedding(**rope_params, device=device) +def _rotate_half(x: Tensor) -> Tensor: + """Rotate the two contiguous halves used by NeoX-style RoPE.""" + x1, x2 = x.chunk(2, dim=-1) + return torch.cat((-x2, x1), dim=-1) + + +@torch.compile(dynamic=True) +def apply_rotary_pos_emb_fp32(query: Tensor, + key: Tensor, + cos: Tensor, + sin: Tensor, + unsqueeze_dim: int = 1) -> tuple[Tensor, Tensor]: + """Apply NeoX-style RoPE with FP32 arithmetic and dtype-preserving output.""" + query_dtype = query.dtype + key_dtype = key.dtype + query = query.float() + key = key.float() + cos = cos.unsqueeze(unsqueeze_dim).float() + sin = sin.unsqueeze(unsqueeze_dim).float() + query = query * cos + _rotate_half(query) * sin + key = key * cos + _rotate_half(key) * sin + return query.to(query_dtype), key.to(key_dtype) + + class ApplyRotaryEmb(nn.Module): """Apply rotary embedding.""" diff --git a/lmdeploy/pytorch/third_party/deep_gemm/__init__.py b/lmdeploy/pytorch/third_party/deep_gemm/__init__.py index cc9cb4333c..336f13634f 100644 --- a/lmdeploy/pytorch/third_party/deep_gemm/__init__.py +++ b/lmdeploy/pytorch/third_party/deep_gemm/__init__.py @@ -78,20 +78,23 @@ def m_grouped_fp8_gemm_nt_contiguous(a, b, d, m_indices, recipe=None, compiled_d try: from deep_gemm import m_grouped_fp8_gemm_nt_masked except Exception: - from deep_gemm import m_grouped_gemm_fp8_fp8_bf16_nt_masked - - def m_grouped_fp8_gemm_nt_masked(a, - b, - d, - masked_m, - expected_m, - recipe=None, - compiled_dims='nk', - disable_ue8m0_cast=False): - assert recipe is None - assert compiled_dims == 'nk' - assert disable_ue8m0_cast is False - return m_grouped_gemm_fp8_fp8_bf16_nt_masked(a, b, d, masked_m, expected_m) + try: + from deep_gemm import fp8_m_grouped_gemm_nt_masked as m_grouped_fp8_gemm_nt_masked + except Exception: + from deep_gemm import m_grouped_gemm_fp8_fp8_bf16_nt_masked + + def m_grouped_fp8_gemm_nt_masked(a, + b, + d, + masked_m, + expected_m, + recipe=None, + compiled_dims='nk', + disable_ue8m0_cast=False): + assert recipe is None + assert compiled_dims == 'nk' + assert disable_ue8m0_cast is False + return m_grouped_gemm_fp8_fp8_bf16_nt_masked(a, b, d, masked_m, expected_m) try: diff --git a/lmdeploy/serve/processors/multimodal.py b/lmdeploy/serve/processors/multimodal.py index ac3a3e7ebb..77ac6ceb3b 100644 --- a/lmdeploy/serve/processors/multimodal.py +++ b/lmdeploy/serve/processors/multimodal.py @@ -54,6 +54,33 @@ def __init__(self, self.backend = backend self.allowed_media_domains = allowed_media_domains + @staticmethod + def merge_media_io_kwargs( + defaults: dict[str, Any] | None, + overrides: dict[str, Any] | None, + ) -> dict[str, Any]: + """Merge model media defaults with request-level modality overrides.""" + merged: dict[str, Any] = {} + for key, value in (defaults or {}).items(): + merged[key] = dict(value) if isinstance(value, dict) else value + for key, value in (overrides or {}).items(): + if isinstance(value, dict) and isinstance(merged.get(key), dict): + merged[key] = {**merged[key], **value} + else: + merged[key] = dict(value) if isinstance(value, dict) else value + return merged + + def resolve_media_io_kwargs( + self, + overrides: dict[str, Any] | None, + ) -> dict[str, Any]: + """Resolve model-specific media defaults without mutating requests.""" + model = getattr(self.vl_encoder, 'model', None) + defaults = getattr(model, 'default_media_io_kwargs', None) + if callable(defaults): + defaults = defaults() + return self.merge_media_io_kwargs(defaults, overrides) + @staticmethod def merge_message_content(msg: dict) -> dict: """Merge multimodal content blocks and ensure content field exists. @@ -411,6 +438,7 @@ async def _get_multimodal_prompt_input(self, """Process multimodal prompt and return processed data for inference engines.""" chat_template = self.chat_template if do_preprocess else BaseChatTemplate() + media_io_kwargs = self.resolve_media_io_kwargs(media_io_kwargs) messages = await self.async_parse_multimodal_item(messages, media_io_kwargs, allowed_media_domains=self.allowed_media_domains) diff --git a/lmdeploy/vl/media/video.py b/lmdeploy/vl/media/video.py index a8c9f2b4ac..c489aa5877 100644 --- a/lmdeploy/vl/media/video.py +++ b/lmdeploy/vl/media/video.py @@ -32,12 +32,17 @@ def __init__( self, image_io: ImageMediaIO, num_frames: int = 32, + max_frames: int | None = None, **kwargs, ) -> None: super().__init__() self.image_io = image_io - self.num_frames = num_frames + # ``num_frames`` is LMDeploy's existing request knob; GLM/SGLang call + # the equivalent upper bound ``max_frames``. Accept both at this media + # boundary so model defaults and per-request overrides share one loader. + self.num_frames = int(max_frames if max_frames is not None else + num_frames) # for potential custom arguments from --media-io-kwargs self.kwargs = kwargs diff --git a/lmdeploy/vl/media/video_loader.py b/lmdeploy/vl/media/video_loader.py index 0cf921f3e5..58f98a6727 100644 --- a/lmdeploy/vl/media/video_loader.py +++ b/lmdeploy/vl/media/video_loader.py @@ -18,6 +18,70 @@ logger = get_logger('lmdeploy') +def glm_sample_frame_indices( + total_frames: int, + source_fps: float, + duration: float, + *, + target_fps: float | None = None, + max_frame_count: int | None = None, +) -> list[int]: + """Sample the deterministic temporal pairs expected by GLM video models. + + GLM constructs one visual unit from every two sampled source frames. Its + serving contract therefore differs from the generic uniform sampler in two + ways: the default is 2 FPS with at most 2048 frames, and an odd result + repeats its final frame to keep temporal pairs complete. + """ + if total_frames <= 0: + return [] + target_fps = 2.0 if target_fps is None else float(target_fps) + max_frame_count = (2048 if max_frame_count is None else + int(max_frame_count)) + if target_fps <= 0 or max_frame_count <= 0: + return [] + + max_frame_idx = total_frames - 1 + if not duration: + duration = (round(max_frame_idx / source_fps) + 1 + if source_fps else 0) + extract_t = min(int(duration * target_fps), max_frame_count) + extract_t = max(1, extract_t) + + if source_fps: + duration_per_frame = 1 / source_fps + max_second = int(duration) + indices = [] + current_second = 0.0 + interval = 1 / target_fps + for frame_index in range(total_frames): + timestamp = frame_index * duration_per_frame + if timestamp >= current_second: + current_second += interval + indices.append(frame_index) + if current_second >= max_second: + break + else: + indices = [] + + if len(indices) < extract_t: + start = indices[0] if indices else 0 + end = indices[-1] if indices else max(total_frames - 1, 0) + indices = np.linspace(start, end, extract_t, dtype=int).tolist() + elif len(indices) > extract_t: + indices = np.linspace(0, + total_frames - 1, + extract_t, + dtype=int).tolist() + + # np.linspace can repeat indices for very small inputs. GLM first removes + # those repeats, then pads an odd number of frames with the final sample. + unique_indices = list(dict.fromkeys(int(index) for index in indices)) + if len(unique_indices) & 1: + unique_indices.append(unique_indices[-1]) + return unique_indices + + class VideoLoader: @classmethod @@ -26,9 +90,27 @@ def load_bytes(self, data: bytes, num_frames: int = -1, **kwargs) -> tuple[npt.N raise NotImplementedError @classmethod - def smart_nframes(self, total_frames_num: int, num_frames: int, fps: int, duration: int) -> tuple[int, list[int]]: + def smart_nframes(self, + total_frames_num: int, + num_frames: int, + fps: float, + duration: float, + sampling_strategy: str = 'uniform', + source_fps: float | None = None) -> tuple[int, list[int]]: # resample video to target num_frames and fps # - the minimum of the two will be used + if sampling_strategy == 'glm': + frame_idx = glm_sample_frame_indices( + total_frames_num, + source_fps=source_fps or 0, + duration=duration, + target_fps=None if fps <= 0 else fps, + max_frame_count=None if num_frames <= 0 else num_frames, + ) + return len(frame_idx), frame_idx + if sampling_strategy != 'uniform': + raise ValueError( + f'Unknown video sampling strategy: {sampling_strategy!r}') num_frames_to_sample = total_frames_num if num_frames > 0: num_frames_to_sample = min(num_frames, total_frames_num) @@ -118,21 +200,28 @@ def load_file( self, filepath: Path, num_frames: int = -1, - fps: int = -1, + fps: float = -1, max_duration: int = 300, + sampling_strategy: str = 'uniform', **kwargs, ) -> tuple[npt.NDArray, dict[str, Any]]: with open(filepath, 'rb') as f: data = f.read() - return self.load_bytes(data, num_frames=num_frames, fps=fps, max_duration=max_duration, **kwargs) + return self.load_bytes(data, + num_frames=num_frames, + fps=fps, + max_duration=max_duration, + sampling_strategy=sampling_strategy, + **kwargs) @classmethod def load_bytes( cls, data: bytes, num_frames: int = -1, - fps: int = -1, + fps: float = -1, max_duration: int = 300, + sampling_strategy: str = 'uniform', **kwargs, ) -> tuple[npt.NDArray, dict[str, Any]]: """Load video frames from bytes. @@ -157,11 +246,35 @@ def load_bytes( original_fps = cap.get(cv2.CAP_PROP_FPS) duration = total_frames_num / original_fps if original_fps > 0 else 0 - num_frames_to_sample, frame_idx = cls.smart_nframes(total_frames_num, num_frames, fps, duration) - - frame_idx_set = set(frame_idx) - frames, valid_num_frames, valid_frame_indices = cls._read_frames(cap, frame_idx_set, num_frames_to_sample, - max(frame_idx)) + _, frame_idx = cls.smart_nframes( + total_frames_num, + num_frames, + fps, + duration, + sampling_strategy=sampling_strategy, + source_fps=original_fps, + ) + if not frame_idx: + raise ValueError('Video sampling produced no frame indices.') + + unique_frame_indices = list(dict.fromkeys(frame_idx)) + frame_idx_set = set(unique_frame_indices) + frames, _, valid_frame_indices = cls._read_frames( + cap, + frame_idx_set, + len(unique_frame_indices), + max(frame_idx), + ) + # The GLM sampler may repeat the final frame to complete a temporal + # pair. OpenCV decodes each source index once, so restore the requested + # order (including repeats) after decoding. + decoded = dict(zip(valid_frame_indices, frames)) + ordered_indices = [index for index in frame_idx if index in decoded] + if ordered_indices: + frames = np.stack([decoded[index] for index in ordered_indices]) + else: + frames = frames[:0] + valid_frame_indices = ordered_indices # Use transformers transformers.video_utils.VideoMetadata format # For models like Qwen3-VL/GLM4.5V, this metadata @@ -186,8 +299,9 @@ class DecordVideoLoader(VideoLoader): def load_file(self, filepath: Path, num_frames: int = -1, - fps: int = -1, + fps: float = -1, max_duration: int = 300, + sampling_strategy: str = 'uniform', **kwargs) -> tuple[npt.NDArray, dict[str, Any]]: import decord vr = decord.VideoReader(str(filepath)) @@ -195,7 +309,16 @@ def load_file(self, original_fps = vr.get_avg_fps() duration = total_frames_num / original_fps if original_fps > 0 else 0 - num_frames_to_sample, frame_idx = self.smart_nframes(total_frames_num, num_frames, fps, duration) + _, frame_idx = self.smart_nframes( + total_frames_num, + num_frames, + fps, + duration, + sampling_strategy=sampling_strategy, + source_fps=original_fps, + ) + if not frame_idx: + raise ValueError('Video sampling produced no frame indices.') video = vr.get_batch(frame_idx).asnumpy() # THWC metadata = { @@ -211,8 +334,9 @@ def load_file(self, def load_bytes(self, data: bytes, num_frames: int = -1, - fps: int = -1, + fps: float = -1, max_duration: int = 300, + sampling_strategy: str = 'uniform', **kwargs) -> tuple[npt.NDArray, dict[str, Any]]: tmp_file = tempfile.NamedTemporaryFile(delete=False, suffix='.mp4') try: @@ -222,6 +346,7 @@ def load_bytes(self, num_frames=num_frames, fps=fps, max_duration=max_duration, + sampling_strategy=sampling_strategy, **kwargs) finally: # always cleanup, even if load_file crashes @@ -237,8 +362,9 @@ class TorchCodecVideoLoader(VideoLoader): def load_file(self, filepath: Path, num_frames: int = -1, - fps: int = -1, + fps: float = -1, max_duration: int = 300, + sampling_strategy: str = 'uniform', **kwargs) -> tuple[npt.NDArray, dict[str, Any]]: # torchcodec requires matched ffmpeg, torchcodec, and torch versions # ffmpeg 5.1.2, torch 2.8.0, torchcodec 0.7.0 are verified to work together @@ -250,7 +376,16 @@ def load_file(self, original_fps = decoder.metadata.average_fps duration = total_frames_num / original_fps if original_fps > 0 else 0 - num_frames_to_sample, frame_idx = self.smart_nframes(total_frames_num, num_frames, fps, duration) + _, frame_idx = self.smart_nframes( + total_frames_num, + num_frames, + fps, + duration, + sampling_strategy=sampling_strategy, + source_fps=original_fps, + ) + if not frame_idx: + raise ValueError('Video sampling produced no frame indices.') video = decoder.get_frames_at(frame_idx).data metadata = { @@ -266,8 +401,9 @@ def load_file(self, def load_bytes(self, data: bytes, num_frames: int = -1, - fps: int = -1, + fps: float = -1, max_duration: int = 300, + sampling_strategy: str = 'uniform', **kwargs) -> tuple[npt.NDArray, dict[str, Any]]: tmp_file = tempfile.NamedTemporaryFile(delete=False, suffix='.mp4') try: @@ -277,6 +413,7 @@ def load_bytes(self, num_frames=num_frames, fps=fps, max_duration=max_duration, + sampling_strategy=sampling_strategy, **kwargs) finally: # always cleanup, even if load_file crashes @@ -292,8 +429,9 @@ class TorchVisionVideoLoader(VideoLoader): def load_file(self, filepath: Path, num_frames: int = -1, - fps: int = -1, + fps: float = -1, max_duration: int = 300, + sampling_strategy: str = 'uniform', **kwargs) -> tuple[npt.NDArray, dict[str, Any]]: import torchvision @@ -306,7 +444,16 @@ def load_file(self, original_fps = info['video_fps'] duration = total_frames_num / original_fps if original_fps > 0 else 0 - num_frames_to_sample, frame_idx = self.smart_nframes(total_frames_num, num_frames, fps, duration) + _, frame_idx = self.smart_nframes( + total_frames_num, + num_frames, + fps, + duration, + sampling_strategy=sampling_strategy, + source_fps=original_fps, + ) + if not frame_idx: + raise ValueError('Video sampling produced no frame indices.') video = video[frame_idx] metadata = { @@ -322,8 +469,9 @@ def load_file(self, def load_bytes(self, data: bytes, num_frames: int = -1, - fps: int = -1, + fps: float = -1, max_duration: int = 300, + sampling_strategy: str = 'uniform', **kwargs) -> tuple[npt.NDArray, dict[str, Any]]: tmp_file = tempfile.NamedTemporaryFile(delete=False, suffix='.mp4') try: @@ -333,6 +481,7 @@ def load_bytes(self, num_frames=num_frames, fps=fps, max_duration=max_duration, + sampling_strategy=sampling_strategy, **kwargs) finally: # always cleanup, even if load_file crashes diff --git a/lmdeploy/vl/model/builder.py b/lmdeploy/vl/model/builder.py index 62aa8c6cb1..a590fa6144 100644 --- a/lmdeploy/vl/model/builder.py +++ b/lmdeploy/vl/model/builder.py @@ -14,6 +14,7 @@ from .gemma3_vl import Gemma3VisionModel # noqa F401 from .glm4_1v import GLM4_1_VisionModel # noqa F401 from .glm4_v import GLM4VisionModel # noqa F401 +from .glm5_next import GLM5NextVisionModel # noqa F401 from .interns1_pro import InternS1ProVisionModel # noqa F401 from .internvl import InternVLVisionModel # noqa F401 from .internvl3_hf import InternVL3VisionModel # noqa F401 diff --git a/lmdeploy/vl/model/glm5_next.py b/lmdeploy/vl/model/glm5_next.py new file mode 100644 index 0000000000..79a3068280 --- /dev/null +++ b/lmdeploy/vl/model/glm5_next.py @@ -0,0 +1,295 @@ +# Copyright (c) OpenMMLab. All rights reserved. +"""Multimodal frontend for GLM-5.3-Flash.""" + +from __future__ import annotations + +import json +from pathlib import Path +from typing import Any + +import torch +from transformers import AutoProcessor, AutoTokenizer + +from lmdeploy.vl.constants import Modality +from lmdeploy.vl.model.base import ( + VISION_MODELS, + MultimodalSpecialTokens, + VisionModel, +) +from lmdeploy.vl.model.preprocess_utils import ( + get_expanded_mm_items, + get_override_size, +) + + +def _processor_kwargs(config: dict[str, Any]) -> dict[str, Any]: + """Translate GLM-5 token budgets to the GLM-4V pixel contract.""" + config = dict(config) + for key in ('image_processor_type', 'video_processor_type', + 'patch_expand_factor'): + config.pop(key, None) + min_tokens = config.pop('min_image_tokens', None) + max_tokens = config.pop('max_image_tokens', None) + patch_size = int(config.get('patch_size', 14)) + merge_size = int(config.get('merge_size', 2)) + temporal_patch_size = int(config.get('temporal_patch_size', 2)) + pixels_per_token = temporal_patch_size * (patch_size * merge_size)**2 + if min_tokens is not None or max_tokens is not None: + config['size'] = { + 'shortest_edge': int(min_tokens or 1) * pixels_per_token, + 'longest_edge': int(max_tokens or min_tokens) * pixels_per_token, + } + return config + + +@VISION_MODELS.register_module() +class GLM5NextVisionModel(VisionModel): + """Prepare GLM-5.3 images/videos for the native PyTorch vision tower.""" + + _arch = ['Glm5NextForConditionalGeneration'] + # Match GLM/SGLang's video contract before the HF processor: sample at + # 2 FPS, cap at 2048 source frames, and complete temporal pairs. + default_media_io_kwargs = { + 'video': { + 'fps': 2.0, + 'num_frames': 2048, + 'sampling_strategy': 'glm', + }, + } + + @classmethod + def match(cls, config): + arch = config.architectures[0] if config.architectures else None + return arch in cls._arch and getattr(config, 'vision_config', None) is not None + + def _build_compat_processor(self, trust_remote_code: bool): + """Build from public GLM-4V components on pre-GLM-5 Transformers.""" + from transformers.models.glm4v.image_processing_glm4v import ( + Glm4vImageProcessor, + ) + from transformers.models.glm4v.processing_glm4v import Glm4vProcessor + from transformers.models.glm4v.video_processing_glm4v import ( + Glm4vVideoProcessor, + ) + + config_path = Path(self.model_path) / 'processor_config.json' + processor_config = json.loads(config_path.read_text()) + image_processor = Glm4vImageProcessor( + **_processor_kwargs(processor_config['image_processor'])) + video_processor = Glm4vVideoProcessor( + **_processor_kwargs(processor_config['video_processor'])) + tokenizer = AutoTokenizer.from_pretrained( + self.model_path, trust_remote_code=trust_remote_code) + return Glm4vProcessor( + image_processor=image_processor, + tokenizer=tokenizer, + video_processor=video_processor, + chat_template=getattr(tokenizer, 'chat_template', None), + ) + + def build_preprocessor(self, trust_remote_code: bool = False): + processor = AutoProcessor.from_pretrained( + self.model_path, trust_remote_code=trust_remote_code) + if not (hasattr(processor, 'image_processor') + and hasattr(processor, 'video_processor')): + processor = self._build_compat_processor(trust_remote_code) + self.processor = processor + + self.image_token = processor.image_token + self.video_token = processor.video_token + self.image_token_id = int(self.hf_config.image_token_id) + configured_video_token_id = int(self.hf_config.video_token_id) + self.input_video_token_id = configured_video_token_id + + # GLM video prompts use <|video|> before processing, then expand each + # frame to an image-token span. Detect this contract rather than + # hard-coding it, so a future native GLM-5 processor can use a distinct + # post-tokenization video id. + frame_builder = getattr(processor, 'replace_frame_token_id', None) + frame_text = frame_builder(0, 1) if callable(frame_builder) else '' + self.video_token_id = (self.image_token_id + if self.image_token in frame_text else + configured_video_token_id) + self._shared_video_token = self.video_token_id == self.image_token_id + self.mm_tokens = MultimodalSpecialTokens( + image_token=self.image_token, + video_token=self.video_token, + image_token_id=self.image_token_id, + video_token_id=self.video_token_id, + ) + + @staticmethod + def _next_span(input_ids: torch.Tensor, cursor: int, token_id: int, + length: int) -> tuple[tuple[int, int], int]: + while cursor < len(input_ids) and int(input_ids[cursor]) != token_id: + cursor += 1 + end = cursor + length + if end > len(input_ids) or not torch.all(input_ids[cursor:end] == token_id): + raise ValueError( + f'cannot locate a contiguous multimodal span of {length} tokens') + return (cursor, end), end + + def _shared_token_offsets( + self, + input_ids: torch.Tensor, + mm_items: list[tuple[Modality, Any, dict]], + collected: dict[Modality, dict[str, Any]], + ) -> None: + """Recover mixed image/video ownership when both use image tokens.""" + merge_length = self.processor.image_processor.merge_size**2 + image_index = 0 + video_index = 0 + cursor = 0 + image_offsets = [] + video_offsets = [] + for modality, _, _ in mm_items: + if modality == Modality.IMAGE: + grid = collected[Modality.IMAGE]['image_grid_thw'][image_index] + length = int(torch.as_tensor(grid).prod().item()) // merge_length + span, cursor = self._next_span(input_ids, cursor, + self.image_token_id, length) + image_offsets.append(span) + image_index += 1 + elif modality == Modality.VIDEO: + grid = collected[Modality.VIDEO]['video_grid_thw'][video_index] + t, h, w = torch.as_tensor(grid).tolist() + length = int(h * w) // merge_length + for _ in range(int(t)): + span, cursor = self._next_span(input_ids, cursor, + self.image_token_id, length) + video_offsets.append(span) + video_index += 1 + if Modality.IMAGE in collected: + collected[Modality.IMAGE]['offset'] = image_offsets + if Modality.VIDEO in collected: + collected[Modality.VIDEO]['offset'] = video_offsets + + def _expand_raw_input_ids( + self, + input_prompt: list[int], + outputs: dict[str, Any], + raw_videos: list[Any], + video_metadatas: list[Any], + ) -> torch.Tensor: + """Expand raw image/video placeholder IDs with official GLM text.""" + tokenizer = self.processor.tokenizer + image_grids = outputs.get('image_grid_thw') + video_grids = outputs.get('video_grid_thw') + + image_replacements: list[list[int]] = [] + if image_grids is not None: + for image_idx in range(len(image_grids)): + replacement = self.processor.replace_image_token( + outputs, image_idx=image_idx) + image_replacements.append( + tokenizer.encode(replacement, add_special_tokens=False)) + + video_replacements: list[list[int]] = [] + if video_grids is not None: + from transformers.video_utils import make_batched_metadata + + metadata = make_batched_metadata( + raw_videos, video_metadata=video_metadatas) + video_inputs = { + 'video_grid_thw': video_grids, + 'video_metadata': metadata, + } + for video_idx in range(len(video_grids)): + replacement = self.processor.replace_video_token( + video_inputs, video_idx=video_idx) + video_replacements.append( + tokenizer.encode(replacement, add_special_tokens=False)) + + image_count = input_prompt.count(self.image_token_id) + video_count = input_prompt.count(self.input_video_token_id) + if image_count != len(image_replacements): + raise ValueError( + 'raw GLM-5 input image placeholders do not match images: ' + f'{image_count} placeholders for {len(image_replacements)} images.' + ) + if video_count != len(video_replacements): + raise ValueError( + 'raw GLM-5 input video placeholders do not match videos: ' + f'{video_count} placeholders for {len(video_replacements)} videos.' + ) + + image_index = 0 + video_index = 0 + expanded: list[int] = [] + for token in input_prompt: + if token == self.image_token_id: + expanded.extend(image_replacements[image_index]) + image_index += 1 + elif token == self.input_video_token_id: + expanded.extend(video_replacements[video_index]) + video_index += 1 + else: + expanded.append(token) + return torch.tensor(expanded, dtype=torch.long) + + def preprocess(self, + messages: list[dict], + input_prompt: str | list[int], + mm_processor_kwargs: dict[str, Any] | None = None): + mm_items = self.collect_multimodal_items(messages) + modalities = {item[0] for item in mm_items} + is_raw_video = (not isinstance(input_prompt, str) + and Modality.VIDEO in modalities) + is_mixed = {Modality.IMAGE, Modality.VIDEO} <= modalities + if not self._shared_video_token or not (is_mixed or is_raw_video): + return super().preprocess(messages, input_prompt, + mm_processor_kwargs) + + raw_images = [item[1] for item in mm_items + if item[0] == Modality.IMAGE] + raw_videos = [item[1] for item in mm_items + if item[0] == Modality.VIDEO] + video_metadatas = [item[2].get('video_metadata') for item in mm_items + if item[0] == Modality.VIDEO] + mm_processor_kwargs = mm_processor_kwargs or {} + kwargs: dict[str, Any] = {} + if raw_images: + kwargs['images'] = raw_images + if raw_videos: + kwargs['videos'] = raw_videos + kwargs['videos_kwargs'] = { + 'video_metadata': video_metadatas, + 'do_resize': True, + 'do_sample_frames': False, + } + image_size = get_override_size(self.processor.image_processor, + mm_processor_kwargs.get('image'), + modality='image') + if image_size is not None: + kwargs['images_kwargs'] = {'size': image_size} + video_size = get_override_size(self.processor.video_processor, + mm_processor_kwargs.get('video'), + modality='video') + if video_size is not None: + kwargs['videos_kwargs']['size'] = video_size + + input_text = input_prompt if isinstance(input_prompt, str) else '' + outputs = self.processor(text=[input_text], + padding=True, + return_tensors='pt', + **kwargs) + collected: dict[Modality, dict[str, Any]] = {} + for name, value in outputs.items(): + modality = self.ATTR_NAME_TO_MODALITY.get(name) + if modality not in (Modality.IMAGE, Modality.VIDEO): + continue + collected.setdefault(modality, {}) + if name in self.FEATURE_NAMES: + value = self._postprocess_mm_output( + value, getattr(self, 'mm_feature_dtype', None)) + name = 'feature' + collected[modality][name] = value + + input_ids = (outputs['input_ids'].flatten() + if isinstance(input_prompt, str) else + self._expand_raw_input_ids(input_prompt, outputs, + raw_videos, + video_metadatas)) + self._shared_token_offsets(input_ids, mm_items, collected) + expanded = get_expanded_mm_items(collected, self.mm_tokens) + return dict(input_ids=input_ids.tolist(), multimodal=expanded) diff --git a/requirements/runtime_cuda.txt b/requirements/runtime_cuda.txt index e29e144312..4db5a4013a 100644 --- a/requirements/runtime_cuda.txt +++ b/requirements/runtime_cuda.txt @@ -3,7 +3,7 @@ accelerate>=0.29.3 aiohttp apache-tvm-ffi==0.1.11; sys_platform == "linux" and "aarch64" not in platform_machine and "arm" not in platform_machine -flash-linear-attention +flash-linear-attention==0.5.2 opencv-python-headless peft<=0.14.0 prometheus_client From fca7e9fc3debcaf84538077dd1475eacfaca518a Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Thu, 17 Sep 2026 05:18:49 +0000 Subject: [PATCH 02/39] refactor(pytorch): use default Triton MoE for GLM-5.3 Flash --- lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py | 14 ++++++++++---- lmdeploy/pytorch/kernels/cuda/moe/blocked_fp8.py | 8 ++++++-- lmdeploy/pytorch/kernels/cuda/moe/fused_moe.py | 9 +++++++-- lmdeploy/pytorch/models/glm5_next.py | 4 ---- 4 files changed, 23 insertions(+), 12 deletions(-) diff --git a/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py b/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py index 3665a88170..3865eb77bb 100644 --- a/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py +++ b/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py @@ -337,13 +337,17 @@ def __init__(self, num_experts: int, renormalize: bool = False, block_size: int = 128, - out_dtype: torch.dtype = torch.float16): + out_dtype: torch.dtype = torch.float16, + fp32_acc: bool = False, + output_scale: float = 1.0): super().__init__() self.num_experts = num_experts self.top_k = top_k self.renormalize = renormalize self.block_size = block_size self.out_dtype = out_dtype + self.fp32_acc = fp32_acc + self.output_scale = output_scale def ep_expert_list(self, world_size: int, rank: int): """Experts list of current rank.""" @@ -392,7 +396,9 @@ def forward(self, expert_offset=expert_offset, num_experts=num_experts, renormalize=self.renormalize, - act_func=act_func) + act_func=act_func, + fp32_acc=self.fp32_acc, + output_scale=self.output_scale) output = output.unflatten(0, input_size[:-1]) return output @@ -626,14 +632,14 @@ def _build_fused_moe_blocked_f8(spec: FusedMoEBlockedF8BuildSpec) -> FusedMoEBlo layer_index=spec.layer_idx, ) else: - if spec.fp32_acc or spec.output_scale != 1.0: - raise ValueError('FP32 MoE reduction and output scaling require use_deep_gemm=True.') impl = TritonFusedMoEBlockedF8Impl( top_k=spec.top_k, num_experts=spec.num_experts, renormalize=spec.renormalize, block_size=spec.block_size, out_dtype=spec.output_dtype, + fp32_acc=spec.fp32_acc, + output_scale=spec.output_scale, ) impl.set_scale_fmt(spec.scale_fmt) return impl diff --git a/lmdeploy/pytorch/kernels/cuda/moe/blocked_fp8.py b/lmdeploy/pytorch/kernels/cuda/moe/blocked_fp8.py index ba0a91f5e2..b12c69b2ab 100644 --- a/lmdeploy/pytorch/kernels/cuda/moe/blocked_fp8.py +++ b/lmdeploy/pytorch/kernels/cuda/moe/blocked_fp8.py @@ -690,7 +690,10 @@ def fused_moe_blocked_fp8(input: torch.Tensor, expert_offset: int = 0, num_experts: int = None, renormalize: bool = False, - act_func: Callable = None) -> torch.Tensor: + act_func: Callable = None, + *, + fp32_acc: bool = False, + output_scale: float = 1.0) -> torch.Tensor: """Fused moe.""" device = input.device M = input.size(0) @@ -816,5 +819,6 @@ def fused_moe_blocked_fp8(input: torch.Tensor, **down_moe_cfg, ) - ret = moe_reduce(intermediate_cache2, topk_weights) + ret = moe_reduce(intermediate_cache2, topk_weights, + fp32_acc=fp32_acc, output_scale=output_scale) return ret diff --git a/lmdeploy/pytorch/kernels/cuda/moe/fused_moe.py b/lmdeploy/pytorch/kernels/cuda/moe/fused_moe.py index 4626694765..ca1d19ed8a 100644 --- a/lmdeploy/pytorch/kernels/cuda/moe/fused_moe.py +++ b/lmdeploy/pytorch/kernels/cuda/moe/fused_moe.py @@ -947,6 +947,7 @@ def _moe_reduce_kernel( N: tl.constexpr, BLOCK_K: tl.constexpr, BLOCK_N: tl.constexpr, + output_scale: tl.constexpr, ): pid = tl.program_id(0) num_n_split = tl.cdiv(N, BLOCK_N) @@ -974,11 +975,14 @@ def _moe_reduce_kernel( wh = h * w[:, None] o = wh.sum(axis=0) + if output_scale != 1.0: + o *= output_scale tl.store(o_ptrs, o, mask=mask_n) -def moe_reduce(hidden_states: torch.Tensor, topk_weights: torch.Tensor, fp32_acc: bool = False) -> torch.Tensor: - """Moe reduce.""" +def moe_reduce(hidden_states: torch.Tensor, topk_weights: torch.Tensor, fp32_acc: bool = False, + *, output_scale: float = 1.0) -> torch.Tensor: + """Weight and reduce experts, optionally scaling before the output cast.""" assert hidden_states.dim() == 3 assert topk_weights.dim() == 2 assert hidden_states.size(0) == topk_weights.size(0) @@ -1008,6 +1012,7 @@ def moe_reduce(hidden_states: torch.Tensor, topk_weights: torch.Tensor, fp32_acc N, BLOCK_K, BLOCK_N, + output_scale, num_warps=num_warps, ) diff --git a/lmdeploy/pytorch/models/glm5_next.py b/lmdeploy/pytorch/models/glm5_next.py index 462cadfc6d..adf8e88ec2 100644 --- a/lmdeploy/pytorch/models/glm5_next.py +++ b/lmdeploy/pytorch/models/glm5_next.py @@ -511,10 +511,6 @@ class Glm5NextMoE(DeepseekV2MoE): # its BF16 store. router_routed_scaling_factor = 1.0 fused_moe_output_scale = 2.5 - # Reuse LMDeploy's generic compact DeepGEMM path. With the model's - # sigmoid KDA gate and single-owner routed scaling restored, its real - # weight replay is closer to the reference than the generic Triton path. - fused_moe_use_deep_gemm = True shared_expert_cls = Glm5NextMLP def __init__(self, config: Any, *args, **kwargs): From d3dbefe003ccd19f514509f8423ad347710059a8 Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Thu, 17 Sep 2026 06:05:11 +0000 Subject: [PATCH 03/39] refactor(pytorch): remove unused EP1 DeepGEMM MoE extensions --- .../pytorch/backends/cuda/moe/blocked_fp8.py | 89 +++---------------- lmdeploy/pytorch/backends/moe.py | 1 - lmdeploy/pytorch/kernels/cuda/moe/ep.py | 11 +-- lmdeploy/pytorch/kernels/cuda/moe/ep_fp8.py | 59 +++--------- lmdeploy/pytorch/models/deepseek_v2.py | 2 - lmdeploy/pytorch/nn/moe/__init__.py | 2 - lmdeploy/pytorch/nn/moe/blocked_fp8.py | 4 +- 7 files changed, 28 insertions(+), 140 deletions(-) diff --git a/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py b/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py index 3865eb77bb..f3bf8ebd2d 100644 --- a/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py +++ b/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py @@ -29,17 +29,7 @@ logger = get_logger('lmdeploy') -def _count_tokens_per_expert(topk_ids: torch.Tensor, - num_experts: int) -> torch.Tensor: - """Count routed assignments with a fixed-size, graph-safe CUDA output.""" - flat_ids = topk_ids.flatten().to(torch.int64) - counts = torch.zeros( - num_experts, dtype=torch.int32, device=topk_ids.device) - return counts.scatter_add_(0, flat_ids, torch.ones_like( - flat_ids, dtype=counts.dtype)) - - -class FusedMoENormal(FusedMoEBlockedF8Impl): +class FusedMoENormal: def __init__( self, @@ -56,34 +46,24 @@ def __init__( num_max_dispatch_tokens_per_rank: int = 128, chunk_size: int | None = 32 * 1024, expert_alignment: int = 128, - renormalize: bool = False, - fp32_acc: bool = False, - output_scale: float = 1.0, ): - super().__init__() self.layer_index = layer_index self.top_k = top_k self.num_experts = num_experts self.block_size = block_size - self.ep_size = ep_size self.num_local_experts = num_experts // ep_size self.out_dtype = out_dtype self.fp8_dtype = fp8_dtype self.scale_fmt = scale_fmt - self.renormalize = renormalize - self.fp32_acc = fp32_acc - self.output_scale = output_scale - self.token_dispatcher = None - if ep_size > 1: - self.token_dispatcher = DeepEPTokenDispatcherNormal( - group=ep_group, - num_experts=num_experts, - num_local_experts=self.num_local_experts, - hidden_size=hidden_dim, - params_dtype=out_dtype, - num_max_dispatch_tokens_per_rank=num_max_dispatch_tokens_per_rank, - expert_alignment=expert_alignment, - ) + self.token_dispatcher = DeepEPTokenDispatcherNormal( + group=ep_group, + num_experts=num_experts, + num_local_experts=self.num_local_experts, + hidden_size=hidden_dim, + params_dtype=out_dtype, + num_max_dispatch_tokens_per_rank=num_max_dispatch_tokens_per_rank, + expert_alignment=expert_alignment, + ) def forward( self, @@ -95,34 +75,7 @@ def forward( down_weights: torch.Tensor, down_scale: torch.Tensor, expert_list: list[int] = None, - gate_up_bias: torch.Tensor = None, - down_bias: torch.Tensor = None, - act_func: Callable = None, ): - if self.token_dispatcher is None: - assert expert_list is None - assert gate_up_bias is None and down_bias is None - input_size = hidden_states.shape - hidden_states = hidden_states.flatten(0, -2) - topk_ids = topk_ids.flatten(0, -2) - topk_weights = _renormalize(topk_weights.flatten(0, -2), self.renormalize) - hs_quant, hs_scale = per_token_group_quant_fp8(hidden_states, - self.block_size, - dtype=up_weights.dtype, - scale_fmt=self.scale_fmt) - tokens_per_expert = _count_tokens_per_expert( - topk_ids, self.num_experts) - out_states = fused_moe_v3_fp8((hs_quant, hs_scale), - topk_ids, - topk_weights, (up_weights, up_scale), - (down_weights, down_scale), - tokens_per_expert, - compact_layout=True, - act_func=act_func, - fp32_acc=self.fp32_acc, - output_scale=self.output_scale) - return out_states.unflatten(0, input_size[:-1]) - hs_quant, hs_scale = per_token_group_quant_fp8(hidden_states, self.block_size, dtype=up_weights.dtype, @@ -138,7 +91,6 @@ def forward( return self.token_dispatcher.combine(out_states) def capture(self): - assert self.token_dispatcher is not None return self.token_dispatcher.buffer_normal.capture() def wait(self, event): @@ -610,27 +562,6 @@ def _build_fused_moe_blocked_f8(spec: FusedMoEBlockedF8BuildSpec) -> FusedMoEBlo num_max_dispatch_tokens_per_rank=spec.num_max_dispatch_tokens_per_rank, layer_idx=spec.layer_idx, ) - elif spec.use_deep_gemm: - try: - import deep_gemm # noqa: F401 - except ImportError as e: - raise ImportError('The DeepGEMM MoE path requires the installable deep_gemm package.') from e - impl = FusedMoENormal( - ep_size=1, - ep_group=spec.ep_group, - num_experts=spec.num_experts, - hidden_dim=spec.hidden_dim, - renormalize=spec.renormalize, - block_size=spec.block_size, - top_k=spec.top_k, - out_dtype=spec.output_dtype, - fp8_dtype=spec.fp8_dtype, - scale_fmt=spec.scale_fmt, - fp32_acc=spec.fp32_acc, - output_scale=spec.output_scale, - num_max_dispatch_tokens_per_rank=spec.num_max_dispatch_tokens_per_rank, - layer_index=spec.layer_idx, - ) else: impl = TritonFusedMoEBlockedF8Impl( top_k=spec.top_k, diff --git a/lmdeploy/pytorch/backends/moe.py b/lmdeploy/pytorch/backends/moe.py index 48b46fa519..800d11ca87 100644 --- a/lmdeploy/pytorch/backends/moe.py +++ b/lmdeploy/pytorch/backends/moe.py @@ -257,7 +257,6 @@ class FusedMoEBlockedF8BuildSpec(BuildSpec[FusedMoEBlockedF8Impl]): scale_fmt: str | None fp32_acc: bool = False output_scale: float = 1.0 - use_deep_gemm: bool = False class FusedMoEV4FP4Impl(ABC): diff --git a/lmdeploy/pytorch/kernels/cuda/moe/ep.py b/lmdeploy/pytorch/kernels/cuda/moe/ep.py index f5d9347c12..ed51f840a8 100644 --- a/lmdeploy/pytorch/kernels/cuda/moe/ep.py +++ b/lmdeploy/pytorch/kernels/cuda/moe/ep.py @@ -142,13 +142,13 @@ def _fwd_kernel_ep_gather( output_tensor_stride1, topk_num: tl.constexpr, BLOCK_D: tl.constexpr, - FP32_ACC: tl.constexpr, - OUTPUT_SCALE: tl.constexpr, ): cur_block = tl.program_id(0) start_cur_token = tl.program_id(1) grid_num = tl.num_programs(1) - compute_dtype = tl.float32 if FP32_ACC else output_tensor.dtype.element_ty + # align with xtuner rl + compute_dtype = output_tensor.dtype.element_ty + # compute_dtype = tl.float32 for cur_token in range(start_cur_token, total_token_num, grid_num): off_d = tl.arange(0, BLOCK_D) @@ -160,7 +160,6 @@ def _fwd_kernel_ep_gather( acc_weight = tl.load(recv_topk_weight + cur_token * recv_topk_weight_stride0 + topk_index) tmp = tl.load(input_tensor + source_token_index * input_tensor_stride0 + cur_block * BLOCK_D + off_d) accumulator += tmp.to(compute_dtype) * acc_weight.to(compute_dtype) - accumulator *= OUTPUT_SCALE tl.store( output_tensor + cur_token * output_tensor_stride0 + cur_block * BLOCK_D + off_d, accumulator.to(output_tensor.dtype.element_ty), @@ -174,8 +173,6 @@ def ep_gather( recv_topk_weight: torch.Tensor, input_index: torch.Tensor, output_tensor: torch.Tensor, - fp32_acc: bool = False, - output_scale: float = 1.0, ): BLOCK_D = 1024 # block size of quantization num_warps = 2 @@ -203,8 +200,6 @@ def ep_gather( topk_num=recv_topk_ids.shape[1], num_warps=num_warps, BLOCK_D=BLOCK_D, - FP32_ACC=fp32_acc, - OUTPUT_SCALE=output_scale, ) return diff --git a/lmdeploy/pytorch/kernels/cuda/moe/ep_fp8.py b/lmdeploy/pytorch/kernels/cuda/moe/ep_fp8.py index 374d69f252..8eab01008f 100644 --- a/lmdeploy/pytorch/kernels/cuda/moe/ep_fp8.py +++ b/lmdeploy/pytorch/kernels/cuda/moe/ep_fp8.py @@ -145,81 +145,50 @@ def _deepgemm_grouped_fp8_nt_contiguous(input_tuple, w_tuple, out: torch.Tensor, return deep_gemm.m_grouped_fp8_gemm_nt_contiguous(input_tuple, w_tuple, out, m_indices) -def _get_compact_all_tokens(num_assignments: int, num_experts: int, block_e: int = 128) -> int: - """Maximum expert-aligned rows for a graph-stable compact layout.""" - max_nonempty_experts = min(num_assignments, num_experts) - return block_e * (max_nonempty_experts + (num_assignments - max_nonempty_experts) // block_e) - - def fused_moe_v3_fp8( hidden_states_fp8: tuple[torch.Tensor, torch.Tensor], topk_idx, topk_weights, w13_weight_fp8: tuple[torch.Tensor, torch.Tensor], w2_weight_fp8: tuple[torch.Tensor, torch.Tensor], - num_recv_tokens_per_expert: list[int] | torch.Tensor | None, - *, - compact_layout: bool = False, - act_func=None, - fp32_acc: bool = False, - output_scale: float = 1.0, + num_recv_tokens_per_expert: list[int] | None, ): hidden_states_fp8, hidden_states_scale = hidden_states_fp8 if num_recv_tokens_per_expert is None: return hidden_states_fp8.to(torch.bfloat16) - if compact_layout: - assert isinstance(num_recv_tokens_per_expert, torch.Tensor) - all_tokens = _get_compact_all_tokens(topk_idx.numel(), num_recv_tokens_per_expert.numel()) - num_recv_tokens_per_expert_gpu = num_recv_tokens_per_expert.to(device=hidden_states_fp8.device, - dtype=torch.int32) - num_recv_tokens_per_expert_gpu = (num_recv_tokens_per_expert_gpu + 127) // 128 * 128 - num_recv_tokens_per_expert_gpu[-1].add_(all_tokens - num_recv_tokens_per_expert_gpu.sum()) - else: - all_tokens = sum(num_recv_tokens_per_expert) - num_recv_tokens_per_expert_gpu = torch.tensor(num_recv_tokens_per_expert, - dtype=torch.int32, - pin_memory=True, - device='cpu').cuda(non_blocking=True) + all_tokens = sum(num_recv_tokens_per_expert) if all_tokens <= 0: return hidden_states_fp8.to(torch.bfloat16) + from lmdeploy.pytorch.third_party.deep_gemm import get_mn_major_tma_aligned_tensor m, k = hidden_states_fp8.size() n = w13_weight_fp8[0].size(1) block_size = k // hidden_states_scale.size(1) gather_out = torch.empty_like(hidden_states_fp8, device=hidden_states_fp8.device, dtype=torch.bfloat16) - # Padding keeps its expert id in the existing scatter kernel. Zero its - # inputs/scales; gather reads only the real routed rows via output_index. - allocator = torch.zeros if compact_layout else torch.empty - input_tensor = allocator((all_tokens, k), device=hidden_states_fp8.device, dtype=hidden_states_fp8.dtype) - input_tensor_scale = allocator((all_tokens, k // block_size), - device=hidden_states_fp8.device, - dtype=torch.float32) + input_tensor = torch.empty((all_tokens, k), device=hidden_states_fp8.device, dtype=hidden_states_fp8.dtype) + input_tensor_scale = torch.empty((all_tokens, k // block_size), + device=hidden_states_fp8.device, + dtype=torch.float32) m_indices = torch.empty(all_tokens, device=hidden_states_fp8.device, dtype=torch.int32) output_index = torch.empty_like(topk_idx) + num_recv_tokens_per_expert_gpu = torch.tensor(num_recv_tokens_per_expert, + dtype=torch.int32, + pin_memory=True, + device='cpu').cuda(non_blocking=True) expert_start_loc = torch.empty_like(num_recv_tokens_per_expert_gpu) ep_scatter_fp8(hidden_states_fp8, hidden_states_scale, topk_idx, num_recv_tokens_per_expert_gpu, expert_start_loc, input_tensor, input_tensor_scale, m_indices, output_index) del hidden_states_fp8 - from lmdeploy.pytorch.third_party.deep_gemm import get_mn_major_tma_aligned_tensor gateup_output = torch.empty((all_tokens, n), device=gather_out.device, dtype=torch.bfloat16) input_tensor_scale = get_mn_major_tma_aligned_tensor(input_tensor_scale) _deepgemm_grouped_fp8_nt_contiguous((input_tensor, input_tensor_scale), w13_weight_fp8, gateup_output, m_indices) - if act_func is None: - down_input = torch.empty((all_tokens, n // 2), device=gateup_output.device, dtype=torch.bfloat16) - silu_and_mul(gateup_output.view(-1, n), down_input) - else: - down_input = act_func(gateup_output.view(-1, n)) + down_input = torch.empty((all_tokens, n // 2), device=gateup_output.device, dtype=torch.bfloat16) + silu_and_mul(gateup_output.view(-1, n), down_input) del gateup_output down_input_fp8, down_input_scale = per_token_group_quant_fp8(down_input, block_size) down_input_scale = get_mn_major_tma_aligned_tensor(down_input_scale) down_output = torch.empty((all_tokens, k), device=gather_out.device, dtype=torch.bfloat16) _deepgemm_grouped_fp8_nt_contiguous((down_input_fp8, down_input_scale), w2_weight_fp8, down_output, m_indices) - ep_gather(down_output, - topk_idx, - topk_weights, - output_index, - gather_out, - fp32_acc=fp32_acc, - output_scale=output_scale) + ep_gather(down_output, topk_idx, topk_weights, output_index, gather_out) return gather_out diff --git a/lmdeploy/pytorch/models/deepseek_v2.py b/lmdeploy/pytorch/models/deepseek_v2.py index 4006bce34b..f8d41d59fd 100644 --- a/lmdeploy/pytorch/models/deepseek_v2.py +++ b/lmdeploy/pytorch/models/deepseek_v2.py @@ -709,7 +709,6 @@ class DeepseekV2MoE(nn.Module): fused_moe_act_func = None fused_moe_fp32_acc = False fused_moe_output_scale = 1.0 - fused_moe_use_deep_gemm = False router_routed_scaling_factor = None shared_expert_cls = None @@ -773,7 +772,6 @@ def __init__(self, act_func=type(self).fused_moe_act_func, fp32_acc=type(self).fused_moe_fp32_acc, output_scale=type(self).fused_moe_output_scale, - use_deep_gemm=type(self).fused_moe_use_deep_gemm, ) self.shared_experts = None if config.n_shared_experts is not None: diff --git a/lmdeploy/pytorch/nn/moe/__init__.py b/lmdeploy/pytorch/nn/moe/__init__.py index fe487a2462..c834f5e32c 100644 --- a/lmdeploy/pytorch/nn/moe/__init__.py +++ b/lmdeploy/pytorch/nn/moe/__init__.py @@ -27,7 +27,6 @@ def build_fused_moe( *, fp32_acc: bool = False, output_scale: float = 1.0, - use_deep_gemm: bool = False, ): """Fused moe builder.""" quant_method = None @@ -113,7 +112,6 @@ def build_fused_moe( act_func=act_func, fp32_acc=fp32_acc, output_scale=output_scale, - use_deep_gemm=use_deep_gemm, ) elif quant_method == 'compressed-tensors': if bias: diff --git a/lmdeploy/pytorch/nn/moe/blocked_fp8.py b/lmdeploy/pytorch/nn/moe/blocked_fp8.py index b85f0dd7e2..fb9c7fd9f0 100644 --- a/lmdeploy/pytorch/nn/moe/blocked_fp8.py +++ b/lmdeploy/pytorch/nn/moe/blocked_fp8.py @@ -159,8 +159,7 @@ def __init__(self, layer_idx: int = 0, act_func: Callable = None, fp32_acc: bool = False, - output_scale: float = 1.0, - use_deep_gemm: bool = False): + output_scale: float = 1.0): device = device or torch.device('cpu') dtype = dtype or torch.float16 @@ -195,7 +194,6 @@ def __init__(self, scale_fmt=scale_fmt, fp32_acc=fp32_acc, output_scale=output_scale, - use_deep_gemm=use_deep_gemm, ), enable_deterministic=build_ctx.enable_deterministic, ) From 6218f599d22eee6007686217982e24830230f61c Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Mon, 21 Sep 2026 07:23:32 +0000 Subject: [PATCH 04/39] feat(pytorch): support GLM-5.3 Flash MTP with shared components --- lmdeploy/pytorch/backends/cuda/kda.py | 62 ++++- lmdeploy/pytorch/configurations/glm5_next.py | 28 +- lmdeploy/pytorch/models/glm5_next.py | 249 ++++++++++++++---- lmdeploy/pytorch/models/glm_moe_dsa_mtp.py | 19 +- lmdeploy/pytorch/models/module_map.py | 1 + .../spec_decode/proposers/deepseek_mtp.py | 3 +- 6 files changed, 299 insertions(+), 63 deletions(-) diff --git a/lmdeploy/pytorch/backends/cuda/kda.py b/lmdeploy/pytorch/backends/cuda/kda.py index 3878eeddea..472ddfbce0 100644 --- a/lmdeploy/pytorch/backends/cuda/kda.py +++ b/lmdeploy/pytorch/backends/cuda/kda.py @@ -6,6 +6,7 @@ semantics and delegates convolution and recurrence to FLA. """ +from copy import copy from typing import Any import torch @@ -59,6 +60,57 @@ def get_step_metadata_provider(self): """Reuse FLA chunk-index preparation outside model forward.""" return GatedDeltaStepMetaUpdater() + def _forward_spec(self, mixed_qkv, raw_gate, raw_beta, conv_state, + recurrent_state, metadata, **kwargs): + """Reuse single-step FLA kernels, checkpointing every verified token. + + State is addressed by accepted history length, not the last proposed + length. The ring therefore also handles zero/partial acceptance and + request reordering without a scheduler-side rollback hook. + """ + history = metadata.cache_seqlens.long() + ring_size = metadata.num_spec_tokens + 1 + ids = metadata.state_ids.long() + read_slot = history.remainder(ring_size) + conv = conv_state[ids, read_slot].clone() + recurrent = recurrent_state[ids, read_slot].clone() + local = copy(metadata) + local.num_spec_tokens = 0 + local.spec_state_offsets = None + local.spec_conv_offsets = None + local.state_ids = torch.arange(ids.numel(), device=ids.device) + + def store(state, values, lengths): + slots = lengths.remainder(ring_size) + previous = state[ids, slots] + valid = metadata.valid_state.reshape(-1, *([1] * (values.ndim - 1))) + state[ids, slots] = torch.where(valid, values.to(state.dtype), previous) + + if not metadata.is_decoding: + output = self.forward(mixed_qkv, raw_gate, raw_beta, + conv_state=conv, recurrent_state=recurrent, + metadata=local, **kwargs) + lengths = history + metadata.cu_seqlens.diff() + store(conv_state, conv, lengths) + store(recurrent_state, recurrent, lengths) + return output + + batch_size = ids.numel() + steps = mixed_qkv.size(1) // batch_size + if steps > ring_size: + raise ValueError('KDA verification exceeds the configured state ring.') + inputs = [x.unflatten(1, (batch_size, steps)) + for x in (mixed_qkv, raw_gate, raw_beta)] + outputs = [] + for step in range(steps): + output = self.forward(*(x[:, :, step].contiguous() for x in inputs), + conv_state=conv, recurrent_state=recurrent, + metadata=local, **kwargs) + store(conv_state, conv, history + step + 1) + store(recurrent_state, recurrent, history + step + 1) + outputs.append(output) + return torch.stack(outputs, dim=2).flatten(1, 2) + def _conv( self, mixed_qkv: torch.Tensor, @@ -112,10 +164,12 @@ def forward( head_dim: int, lower_bound: float, ) -> torch.Tensor: - if (metadata.spec_state_offsets is not None - or (getattr(metadata, 'num_spec_tokens', 0) or 0) > 0): - raise NotImplementedError( - 'GLM-5.3 KDA speculative state rollback is not implemented.') + if (getattr(metadata, 'num_spec_tokens', 0) or 0) > 0: + return self._forward_spec( + mixed_qkv, raw_gate, raw_beta, conv_state, recurrent_state, + metadata, conv_weight=conv_weight, conv_bias=conv_bias, + a_log=a_log, dt_bias=dt_bias, num_heads=num_heads, + head_dim=head_dim, lower_bound=lower_bound) batch_size = metadata.state_ids.numel() if metadata.is_decoding: query_length = mixed_qkv.size(1) // batch_size diff --git a/lmdeploy/pytorch/configurations/glm5_next.py b/lmdeploy/pytorch/configurations/glm5_next.py index 2ad5b5b152..2ef73daf4a 100644 --- a/lmdeploy/pytorch/configurations/glm5_next.py +++ b/lmdeploy/pytorch/configurations/glm5_next.py @@ -213,8 +213,13 @@ def build(cls, hf_config, model_path: str | None = None, **kwargs): linear_config, linear_layer_ids, full_attention_layer_ids = ( _resolve_glm5_linear_config(text_config)) + is_draft = kwargs.get('is_draft_model', False) + num_spec_tokens = kwargs.get('num_spec_tokens', 0) + if is_draft and getattr(text_config, 'num_nextn_predict_layers', 0) != 1: + raise ValueError('GLM-5.3 MTP requires one checkpoint predictor layer.') config = DeepseekV32ModelConfigBuilder.build( - text_config, model_path=model_path, **kwargs) + text_config, model_path=model_path, + **dict(kwargs, is_draft_model=False)) tp = kwargs.get('tp', 1) device_type = kwargs.get('device_type', 'auto') @@ -239,26 +244,29 @@ def build(cls, hf_config, model_path: str | None = None, **kwargs): # unset also preserves the BF16 latent MLA cache policy. config.mla_index_topk = None config.k_head_dim = text_config.kv_lora_rank + 64 + # Keep a complete state after each verified token. Accepted sequence + # lengths select the correct ring slot after rejection sampling. + ring_shape = (num_spec_tokens + 1,) if num_spec_tokens else () config.state_cache_specs = [ StateCacheSpec( GLM5_KDA_CONV_STATE, - (num_linear_layers, conv_dim, conv_kernel_size), + (num_linear_layers, *ring_shape, conv_dim, conv_kernel_size), torch.bfloat16, ), StateCacheSpec( GLM5_KDA_RECURRENT_STATE, - (num_linear_layers, local_heads, head_dim, head_dim), + (num_linear_layers, *ring_shape, local_heads, head_dim, head_dim), torch.float32, ), StateCacheSpec( GLM5_KPOOL_TAIL_K_STATE, - (num_full_layers, text_config.index_kpool, + (num_full_layers, *ring_shape, text_config.index_kpool, text_config.index_head_dim), torch.bfloat16, ), StateCacheSpec( GLM5_KPOOL_TAIL_SCORE_STATE, - (num_full_layers, text_config.index_kpool, + (num_full_layers, *ring_shape, text_config.index_kpool, text_config.index_head_dim), torch.bfloat16, ), @@ -276,6 +284,16 @@ def build(cls, hf_config, model_path: str | None = None, **kwargs): config.check_env_func = _check_env_glm5_next config.hf_config = hf_config config.llm_config = text_config + if is_draft: + # The predictor has MLA/KPool but no KDA or mHC. Its unfinished + # KPool tail is reconstructed from its own pageable token cache. + hf_config.architectures = ['Glm5NextMTPModel'] + if hasattr(hf_config, 'auto_map'): + del hf_config.auto_map + config.num_layers = 1 + config.state_cache_specs = [] + config.states_shapes = [] + config.is_gated_delta = False text_dtype = getattr(text_config, 'dtype', None) if text_dtype is not None: diff --git a/lmdeploy/pytorch/models/glm5_next.py b/lmdeploy/pytorch/models/glm5_next.py index d935f81583..886e50850a 100644 --- a/lmdeploy/pytorch/models/glm5_next.py +++ b/lmdeploy/pytorch/models/glm5_next.py @@ -30,7 +30,8 @@ GLM5_KPOOL_TAIL_SCORE_STATE, ) from lmdeploy.pytorch.distributed import get_dist_manager, get_tp_world_rank -from lmdeploy.pytorch.model_inputs import StepContext, StepContextManager +from lmdeploy.pytorch.engine.cache_engine.schema import BlockCacheRequest +from lmdeploy.pytorch.model_inputs import StepContext, StepContextManager, get_step_ctx_manager from lmdeploy.pytorch.nn import ( FlashAttention, FP32LayerNorm, @@ -69,6 +70,7 @@ Glm4vVisionPatchEmbed, Glm4vVisionRotaryEmbedding, ) +from .glm_moe_dsa_mtp import GlmMoeDsaMTPModel, GlmMoeDsaMultiTokenPredictor from .qwen3_vl import Qwen3VLInputProcessor from .utils.model import build_embedding, vlm_model @@ -513,8 +515,9 @@ class Glm5NextMoE(DeepseekV2MoE): fused_moe_output_scale = 2.5 shared_expert_cls = Glm5NextMLP - def __init__(self, config: Any, *args, **kwargs): - super().__init__(config, *args, **kwargs) + def __init__(self, config: Any, layer_idx: int, *args, **kwargs): + kwargs.setdefault('prefix', f'model.layers.{layer_idx}.mlp') + super().__init__(config, layer_idx, *args, **kwargs) if self.gate.fake_eplb or self.gate.eplb_dispatch_info is not None: raise RuntimeError( 'The GLM-5.3 router does not permit fake ' @@ -741,7 +744,8 @@ def __init__(self, layer_idx, dtype=dtype, device=device, - all_reduce=all_reduce) + all_reduce=all_reduce, + prefix=f'model.layers.{layer_idx}.self_attn') # DeepSeek keeps these latent-norm parameters in FP32. GLM-5.3's # checkpoint and SGLang runtime keep them in the activation dtype; # rebuild only these two containers before weight loading. @@ -753,8 +757,11 @@ def __init__(self, self.index_topk = config.index_topk self.index_kpool = config.index_kpool try: - self.cache_layer_idx = config.full_attention_layer_ids.index( - layer_idx) + full_layer_ids = list(config.full_attention_layer_ids) + full_layer_ids.extend(range(config.num_hidden_layers, + config.num_hidden_layers + + getattr(config, 'num_nextn_predict_layers', 0))) + self.cache_layer_idx = full_layer_ids.index(layer_idx) except ValueError as error: raise ValueError( f'GLM-5.3 full-attention layer {layer_idx} is missing from ' @@ -851,44 +858,58 @@ def _update_kpool_cache( raise RuntimeError('GLM-5.3 KPool requires stable state cache ids.') tail_k_state, tail_score_state = tail_state + history_lengths = attn_metadata.kv_seqlens - attn_metadata.q_seqlens + ring_states = None + if tail_k_state.ndim == 4: + ring_states = tail_state + ring_size = tail_k_state.size(1) + request_ids = state_ids.clamp_min(0).long() + valid_requests = state_ids >= 0 + read_slots = history_lengths.long().remainder(ring_size) + tail_k_state = tail_k_state[request_ids, read_slots].clone() + tail_score_state = tail_score_state[request_ids, read_slots].clone() + local_ids = torch.arange(state_ids.numel(), device=state_ids.device) + # Padding uses its own scratch row, never live request row zero. + state_ids = local_ids + + def save_ring(lengths): + if ring_states is None: + return + slots = lengths.long().remainder(ring_size) + for state, value in zip(ring_states, (tail_k_state, tail_score_state)): + previous = state[request_ids, slots] + state[request_ids, slots] = torch.where( + valid_requests[:, None, None], value, previous) + indexer_k_cache = self.indexer.get_block_cache() key = self.indexer.project_key(hidden_states)[0] score = self.indexer.project_compress_score(hidden_states)[0] if attn_metadata.is_decoding: - if key.size(0) != state_ids.numel(): + batch_size = state_ids.numel() + if key.size(0) % batch_size: raise RuntimeError( - 'KPool decode requires one token per request state.') - history_lengths = (attn_metadata.kv_seqlens - - attn_metadata.q_seqlens) - update = kpool_decode_update( - key, - score, - tail_k_state, - tail_score_state, - state_ids, - history_lengths, - self.index_kpool, - ) - pooled_fp8, pooled_scale = kpool_compress_quantize_cuda( - update.closed_keys, - update.closed_scores, - self.indexer.index_kpool_compress_ape, - mode='decode', - round_scale=self.indexer.scale_fmt is not None, - ) - kpool_write_packed_cache_batched( - indexer_k_cache, - attn_metadata.block_offsets, - update.group_ids, - pooled_fp8, - pooled_scale, - self.index_kpool, - update.should_close, - ) - tail_k_state.index_copy_(0, update.safe_state_ids, - update.next_tail_keys) - tail_score_state.index_copy_(0, update.safe_state_ids, - update.next_tail_scores) + 'KPool decode rows must be divisible by request count.') + steps = key.size(0) // batch_size + key = key.unflatten(0, (batch_size, steps)) + score = score.unflatten(0, (batch_size, steps)) + for step in range(steps): + update = kpool_decode_update( + key[:, step], score[:, step], tail_k_state, + tail_score_state, state_ids, history_lengths + step, + self.index_kpool) + pooled_fp8, pooled_scale = kpool_compress_quantize_cuda( + update.closed_keys, update.closed_scores, + self.indexer.index_kpool_compress_ape, + mode='decode', round_scale=self.indexer.scale_fmt is not None) + kpool_write_packed_cache_batched( + indexer_k_cache, attn_metadata.block_offsets, + update.group_ids, pooled_fp8, pooled_scale, + self.index_kpool, update.should_close) + tail_k_state.index_copy_(0, update.safe_state_ids, + update.next_tail_keys) + tail_score_state.index_copy_(0, update.safe_state_ids, + update.next_tail_scores) + save_ring(history_lengths + step + 1) return indexer_k_cache q_seqlens = attn_metadata.q_seqlens.tolist() @@ -963,6 +984,7 @@ def _update_kpool_cache( raise RuntimeError( f'KPool metadata accounts for {token_offset} tokens, ' f'but projections contain {key.size(0)}.') + save_ring(attn_metadata.kv_seqlens) return indexer_k_cache def _select_kpool_indices( @@ -987,14 +1009,18 @@ def _select_kpool_indices( query_weight = (head_gate * query_scale.squeeze(-1) * self.indexer.softmax_scale) if attn_metadata.is_decoding: - seq_lens = attn_metadata.kv_seqlens.to(torch.int64) + batch_size = attn_metadata.kv_seqlens.numel() + steps = total_rows // batch_size + history = attn_metadata.kv_seqlens - attn_metadata.q_seqlens + step_ids = torch.arange(1, steps + 1, device=query_fp8.device) + seq_lens = (history[:, None] + step_ids).flatten().to(torch.int64) group_lengths = torch.div( seq_lens, self.index_kpool, rounding_mode='floor', ) pooled_block_offsets = kpool_pooled_block_offsets( - attn_metadata.block_offsets, + attn_metadata.block_offsets.repeat_interleave(steps, dim=0), self.index_kpool, ) logits = kpool_score_paged_cuda( @@ -1288,7 +1314,8 @@ def __init__(self, and layer_idx % config.moe_layer_freq == 0) self.mlp = (Glm5NextMoE(config, layer_idx, dtype=dtype, device=device) if is_sparse else Glm5NextMLP( - config, dtype=dtype, device=device)) + config, dtype=dtype, device=device, + prefix=f'model.layers.{layer_idx}.mlp')) self.input_layernorm = RMSNorm(config.hidden_size, config.rms_norm_eps, @@ -1498,6 +1525,7 @@ def forward( vision_groups: list[dict[str, Any]] | None = None, vision_prompt_order: list[int] | None = None, multimodal_mask: torch.Tensor | None = None, + return_input_embeds: bool = False, **kwargs, ) -> torch.Tensor: if inputs_embeds is None and (vision_groups or @@ -1533,13 +1561,19 @@ def forward( inputs_embeds) inputs_embeds = inputs_embeds.masked_scatter( scatter_mask, vision_embeddings.to(inputs_embeds)) - return self.model(input_ids=input_ids, + if return_input_embeds and inputs_embeds is None: + inputs_embeds = self.get_input_embeddings()(input_ids) + hidden_states = self.model(input_ids=input_ids, position_ids=position_ids, past_key_values=past_key_values, attn_metadata=attn_metadata, inputs_embeds=inputs_embeds, state_ids=state_ids, kpool_tail_states=kpool_tail_states) + if return_input_embeds: + return dict(hidden_states=hidden_states, + target_inputs_embeds=inputs_embeds) + return hidden_states def get_input_embeddings(self): return self.model.get_input_embeddings() @@ -1739,14 +1773,18 @@ def prepare_inputs_for_generation( grid_thw=grid_thw, vision_groups=vision_groups, vision_prompt_order=vision_prompt_order, - multimodal_mask=multimodal_mask) + multimodal_mask=multimodal_mask, + return_input_embeds=( + self.ctx_mgr.build_ctx.num_spec_tokens > 0 + and not context.is_decoding)) @staticmethod def _layer_idx(name: str) -> int | None: match = re.search(r'\.layers\.(\d+)\.', name) return None if match is None else int(match.group(1)) - def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]): + def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]], *, + is_mtp: bool = False): """Load both towers through LMDeploy's TP-aware weight loaders.""" stacked_params_mapping = [ ('.gate_up_proj', '.gate_proj', 0), @@ -1782,8 +1820,12 @@ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]): name = checkpoint_name.replace('model.language_model.', 'model.', 1) layer_idx = self._layer_idx(name) - if layer_idx is not None and layer_idx >= self.config.num_hidden_layers: - # MTP layer 45 is a separate milestone. + if is_mtp: + if layer_idx != self.config.num_hidden_layers: + continue + name = self._rewrite_spec_layer_name(layer_idx, name) + elif layer_idx is not None and layer_idx >= self.config.num_hidden_layers: + # Predictor weights are loaded by Glm5NextMTPModel. continue if 'rotary_emb.' in name: continue @@ -1841,3 +1883,116 @@ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]): break else: load_weight(params_dict[name], loaded_weight) + + +class Glm5NextMTPAttention(Glm5NextSparseAttention): + """The predictor's KPool tail is reconstructed from pageable token data. + + Draft forwards can revisit accepted positions after multiple proposals. + Keeping raw index keys/scores in its cache avoids private mutable request + state and reuses the normal cache allocation, sizing and sleep lifecycle. + Only the single MTP layer requests this additional cache. + """ + + _TOKEN_CACHE = 'glm5_mtp_kpool_tokens' + + def get_block_cache_requests(self, context): + return (BlockCacheRequest( + name=self._TOKEN_CACHE, + shape=(context.geometry.kernel_block_size, 2, self.indexer.head_dim), + dtype=torch.bfloat16, + per_row_contiguous=True),) + + def bind_block_cache(self, binding): + if binding.cache_name != self._TOKEN_CACHE: + raise ValueError(f'Unexpected MTP token cache: {binding.cache_name}') + self._token_cache_binding = binding + + def _update_kpool_cache(self, hidden_states, tail_state, state_ids, + attn_metadata): + binding = self._token_cache_binding + caches = get_step_ctx_manager().current_context().block_caches + cache = (caches.row(binding.cache_name, binding.consumer_row) + if hasattr(caches, 'row') else + caches[binding.cache_name][binding.consumer_row]) + block_size = cache.size(1) + history = (attn_metadata.kv_seqlens - attn_metadata.q_seqlens).long() + tail_length = history.remainder(self.index_kpool) + slots = torch.arange(self.index_kpool, device=history.device) + positions = history[:, None] - tail_length[:, None] + slots + block_offsets = attn_metadata.block_offsets.long() + blocks = block_offsets.gather(1, positions.div(block_size, rounding_mode='floor')) + tails = cache[blocks, positions.remainder(block_size)] + tails = tails.masked_fill((slots >= tail_length[:, None])[..., None, None], 0) + state_ids = torch.arange(history.numel(), device=history.device) + result = super()._update_kpool_cache( + hidden_states, (tails[:, :, 0].contiguous(), tails[:, :, 1].contiguous()), + state_ids, attn_metadata) + + # Write raw projected tokens after reading the pre-forward tail. + # Rejected positions are overwritten on their next visit. + total_tokens = hidden_states.size(1) + batch = state_ids.repeat_interleave(attn_metadata.q_seqlens, + output_size=total_tokens) + token_ids = torch.arange(total_tokens, device=history.device) + positions = history[batch] + token_ids - attn_metadata.cu_seqlens_q[batch] + blocks = block_offsets[batch, positions.div(block_size, rounding_mode='floor')] + key = self.indexer.project_key(hidden_states)[0] + score = self.indexer.project_compress_score(hidden_states)[0] + cache[blocks, positions.remainder(block_size)] = torch.stack((key, score), dim=1) + return result + + +class Glm5NextMTPDecoderLayer(nn.Module): + """Checkpoint predictor block: GLM MLA/MoE with plain residuals, no mHC.""" + + def __init__(self, config, layer_idx, dtype=None, device=None): + super().__init__() + self.self_attn = Glm5NextMTPAttention(config, layer_idx, dtype=dtype, device=device) + self.mlp = Glm5NextMoE(config, layer_idx, dtype=dtype, device=device) + self.input_layernorm = RMSNorm(config.hidden_size, config.rms_norm_eps, + dtype=dtype, device=device) + self.post_attention_layernorm = RMSNorm(config.hidden_size, config.rms_norm_eps, + dtype=dtype, device=device) + + def forward(self, hidden_states, rotary_pos_emb, past_key_value, + attn_metadata=None, **kwargs): + residual = hidden_states + hidden_states = self.input_layernorm(hidden_states) + hidden_states = self.self_attn(hidden_states, past_key_value, attn_metadata) + hidden_states, residual = self.post_attention_layernorm(hidden_states, residual) + return self.mlp(hidden_states), residual + + +class Glm5NextMTPModel(GlmMoeDsaMTPModel): + """Reuse the shared GLM/DeepSeek predictor, proposer and CUDA Graph flow.""" + + uses_shared_input_embeddings = True + + def __init__(self, config, ctx_mgr, dtype=None, device=None): + nn.Module.__init__(self) + self.config = config.text_config + self.quantization_config = getattr(config, 'quantization_config', None) + if self.quantization_config is not None: + self.config.quantization_config = self.quantization_config + self.dtype = dtype + self.ctx_mgr = ctx_mgr + self.model = GlmMoeDsaMultiTokenPredictor( + self.config, dtype=dtype, device=device, + decoder_layer_cls=Glm5NextMTPDecoderLayer) + self.uses_dsa_topk_buffer = False + self.topk_indices_buffer = None + self._load_buffers = {} + + def prepare_inputs_for_generation(self, past_key_values, inputs_embeds=None, + context=None): + if context.target_inputs_embeds is not None: + inputs_embeds = context.target_inputs_embeds + return super().prepare_inputs_for_generation(past_key_values, inputs_embeds, context) + + _layer_idx = staticmethod(Glm5NextForConditionalGeneration._layer_idx) + + def load_weights(self, weights): + weights = ((name, weight) for name, weight in weights + if self._layer_idx(name) == self.config.num_hidden_layers) + Glm5NextForConditionalGeneration.load_weights(self, weights, is_mtp=True) diff --git a/lmdeploy/pytorch/models/glm_moe_dsa_mtp.py b/lmdeploy/pytorch/models/glm_moe_dsa_mtp.py index 754a301528..fd5ed7dd77 100644 --- a/lmdeploy/pytorch/models/glm_moe_dsa_mtp.py +++ b/lmdeploy/pytorch/models/glm_moe_dsa_mtp.py @@ -44,6 +44,7 @@ def __init__( layer_idx: int, dtype: torch.dtype = None, device: torch.device = None, + decoder_layer_cls=GlmMoeDsaDecoderLayer, ) -> None: super().__init__() self.enorm = RMSNorm(config.hidden_size, @@ -67,11 +68,12 @@ def __init__( self.shared_head = GlmMoeDsaSharedHead(config, dtype=dtype, device=device) - self.mtp_block = GlmMoeDsaDecoderLayer(config, + self.mtp_block = decoder_layer_cls(config, layer_idx=layer_idx, dtype=dtype, device=device) - self.rotary_emb = build_deepseek_rotary_embedding(config) + self.rotary_emb = (build_deepseek_rotary_embedding(config) + if config.qk_rope_head_dim else None) def forward( self, @@ -89,10 +91,13 @@ def forward( hidden_states = self.eh_proj( torch.cat([inputs_embeds, previous_hidden_states], dim=-1)) - cos, sin = self.rotary_emb(hidden_states, position_ids) + rotary_pos_emb = None + if self.rotary_emb is not None: + cos, sin = self.rotary_emb(hidden_states, position_ids) + rotary_pos_emb = (cos[0], sin[0]) hidden_states, residual = self.mtp_block( hidden_states, - (cos[0], sin[0]), + rotary_pos_emb, past_key_value, attn_metadata=attn_metadata, topk_indices_buffer=topk_indices_buffer, @@ -107,7 +112,8 @@ class GlmMoeDsaMultiTokenPredictor(nn.Module): def __init__(self, config: PretrainedConfig, dtype: torch.dtype = None, - device: torch.device = None): + device: torch.device = None, + decoder_layer_cls=GlmMoeDsaDecoderLayer): super().__init__() self.config = config self.mtp_start_layer_idx = config.num_hidden_layers @@ -118,7 +124,8 @@ def __init__(self, GlmMoeDsaMultiTokenPredictorLayer(config, idx, dtype=dtype, - device=device) + device=device, + decoder_layer_cls=decoder_layer_cls) for idx in range(self.mtp_start_layer_idx, self.mtp_start_layer_idx + self.num_mtp_layers) }) diff --git a/lmdeploy/pytorch/models/module_map.py b/lmdeploy/pytorch/models/module_map.py index b4c09fa4d8..6e6dae21b8 100644 --- a/lmdeploy/pytorch/models/module_map.py +++ b/lmdeploy/pytorch/models/module_map.py @@ -67,6 +67,7 @@ MODULE_MAP.update({ 'Glm5NextForConditionalGeneration': f'{LMDEPLOY_PYTORCH_MODEL_PATH}.glm5_next.Glm5NextForConditionalGeneration', + 'Glm5NextMTPModel': f'{LMDEPLOY_PYTORCH_MODEL_PATH}.glm5_next.Glm5NextMTPModel', }) # internlm2 diff --git a/lmdeploy/pytorch/spec_decode/proposers/deepseek_mtp.py b/lmdeploy/pytorch/spec_decode/proposers/deepseek_mtp.py index 3e2effe361..79135cab86 100644 --- a/lmdeploy/pytorch/spec_decode/proposers/deepseek_mtp.py +++ b/lmdeploy/pytorch/spec_decode/proposers/deepseek_mtp.py @@ -15,7 +15,8 @@ def build_model(self, empty_init: bool, target_model: torch.nn.Module = None, bu super().build_model(empty_init, target_model=target_model, build_model_ctx=build_model_ctx) draft_model = self.model if (getattr(draft_model, 'uses_dsa_topk_buffer', False) - and hasattr(draft_model, 'set_input_embeddings')): + or getattr(draft_model, 'uses_shared_input_embeddings', False)) and hasattr( + draft_model, 'set_input_embeddings'): draft_model.set_input_embeddings(target_model.get_input_embeddings()) async def get_outputs(self, From 3fcabb71ad5e03648aaeef2f28e54271c0524a3c Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Tue, 22 Sep 2026 07:26:51 +0000 Subject: [PATCH 05/39] fix(pytorch): correct GLM-5.3 MTP verification and acceptance Reuse GDN state helpers and channelwise TileLang verification; honor MTP index sharing and invalidate graphs after buffer growth. Share AR sampling filters with rejection/recovery and stabilize GLM TP/mHC and sparse-index arithmetic. Keep EP/DP, prefix-cache and performance experiments outside this patch. --- .../pytorch/backends/cuda/gated_delta_rule.py | 20 ++- lmdeploy/pytorch/backends/cuda/kda.py | 132 +++++++++++------- lmdeploy/pytorch/backends/cuda/kpool.py | 3 + lmdeploy/pytorch/engine/logits_process.py | 73 +++++----- .../pytorch/kernels/cuda/gated_delta_rule.py | 34 +++-- .../pytorch/kernels/cuda/sparse_index_topk.py | 122 ++++++++++------ lmdeploy/pytorch/models/glm5_next.py | 60 ++++++-- lmdeploy/pytorch/models/glm_moe_dsa_mtp.py | 10 +- lmdeploy/pytorch/nn/hc_prepost.py | 11 +- lmdeploy/pytorch/nn/linear/base.py | 12 +- .../pytorch/spec_decode/reject_sampler.py | 6 +- 11 files changed, 327 insertions(+), 156 deletions(-) diff --git a/lmdeploy/pytorch/backends/cuda/gated_delta_rule.py b/lmdeploy/pytorch/backends/cuda/gated_delta_rule.py index 7e07d3b5d0..d3b228cb12 100644 --- a/lmdeploy/pytorch/backends/cuda/gated_delta_rule.py +++ b/lmdeploy/pytorch/backends/cuda/gated_delta_rule.py @@ -150,6 +150,8 @@ def _state_select_kernel( stride_o0, INNER_SIZE: tl.constexpr, BLOCK_SIZE: tl.constexpr, + NUM_STATES: tl.constexpr, + NUM_SLOTS: tl.constexpr, ): """Fused state select: out[b] = state[state_indices[b], spec_offsets[b]].""" @@ -165,7 +167,8 @@ def _state_select_kernel( src_ptr = state_ptr + state_idx * stride_s0 + spec_off * stride_s1 + offs dst_ptr = out_ptr + batch_id * stride_o0 + offs - data = tl.load(src_ptr, mask=mask) + valid = (state_idx >= 0) & (state_idx < NUM_STATES) & (spec_off >= 0) & (spec_off < NUM_SLOTS) + data = tl.load(src_ptr, mask=mask & valid, other=0) tl.store(dst_ptr, data, mask=mask) @@ -180,6 +183,8 @@ def _state_scatter_kernel( stride_i0, INNER_SIZE: tl.constexpr, BLOCK_SIZE: tl.constexpr, + NUM_STATES: tl.constexpr, + NUM_SLOTS: tl.constexpr, ): """Fused state scatter: state[si[b], so[b]] = src[b].""" batch_id = tl.program_id(0).to(tl.int64) @@ -189,7 +194,8 @@ def _state_scatter_kernel( spec_off = tl.load(spec_offsets_ptr + batch_id) offs = block_id * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) - mask = offs < INNER_SIZE + mask = (offs < INNER_SIZE) & (state_idx >= 0) & (state_idx < NUM_STATES) + mask = mask & (spec_off >= 0) & (spec_off < NUM_SLOTS) in_ptr = src_ptr + batch_id * stride_i0 + offs dst_ptr = state_ptr + state_idx * stride_s0 + spec_off * stride_s1 + offs @@ -201,7 +207,8 @@ def _state_scatter_kernel( def _state_select(state, state_indices, spec_offsets): """Fused state select: out = state[state_indices, spec_offsets]. - Requires inner dims [2:] to be contiguous. + Requires inner dims [2:] to be contiguous. Invalid state/slot ids yield + zeros, allowing callers to clear initial or padded requests in this load. """ B = state_indices.shape[0] inner_shape = state.shape[2:] @@ -225,6 +232,8 @@ def _state_select(state, state_indices, spec_offsets): out.stride(0), INNER_SIZE=inner_size, BLOCK_SIZE=BLOCK_SIZE, + NUM_STATES=state.shape[0], + NUM_SLOTS=state.shape[1], ) return out @@ -233,7 +242,8 @@ def _state_scatter(state, state_indices, spec_offsets, src): """Fused state scatter: state[state_indices, spec_offsets] = src.to(state.dtype). - Requires inner dims [2:] to be contiguous. + Requires inner dims [2:] to be contiguous. Invalid state/slot ids are + ignored so padded requests cannot overwrite a live cache row. """ if src.dtype != state.dtype: src = src.to(state.dtype) @@ -257,6 +267,8 @@ def _state_scatter(state, state_indices, spec_offsets, src): src.stride(0), INNER_SIZE=inner_size, BLOCK_SIZE=BLOCK_SIZE, + NUM_STATES=state.shape[0], + NUM_SLOTS=state.shape[1], ) diff --git a/lmdeploy/pytorch/backends/cuda/kda.py b/lmdeploy/pytorch/backends/cuda/kda.py index 472ddfbce0..aaef5a2180 100644 --- a/lmdeploy/pytorch/backends/cuda/kda.py +++ b/lmdeploy/pytorch/backends/cuda/kda.py @@ -1,9 +1,10 @@ # Copyright (c) OpenMMLab. All rights reserved. -"""CUDA KDA backend composed from the public FLA operators. +"""CUDA KDA backend composed from FLA and shared LMDeploy operators. KDA is distinct from LMDeploy's gated-delta rule, but its CUDA implementation does not need copied GLM kernels. This adapter owns only LMDeploy cache/state -semantics and delegates convolution and recurrence to FLA. +semantics, reuses FLA convolution/prefill, and shares the TileLang recurrent +state-ring kernel with gated-delta rule for both AR and MTP decode. """ from copy import copy @@ -13,31 +14,28 @@ from lmdeploy.pytorch.backends.kda import KdaImpl -from .gated_delta_rule import GatedDeltaStepMetaUpdater +from .gated_delta_rule import GatedDeltaStepMetaUpdater, _state_scatter, _state_select from .step_metadata import register_step_metadata_impl def _select_state(state: torch.Tensor, metadata: Any) -> torch.Tensor: - selected = state.index_select(0, metadata.state_ids.long()) - clear = ~metadata.valid_state + valid = metadata.valid_state if metadata.is_init is not None: - clear = clear | metadata.is_init - clear = clear.reshape(-1, *((1, ) * (state.ndim - 1))) - return selected.masked_fill(clear, 0) + valid = valid & ~metadata.is_init + ids = torch.where(valid, metadata.state_ids, -1) + # Treat the ordinary AR bank as a one-slot ring; share the same masked + # gather/scatter kernels with GDN and speculative state checkpoints. + return _state_select(state.unsqueeze(1), ids, torch.zeros_like(ids)) def _store_state(state: torch.Tensor, value: torch.Tensor, metadata: Any) -> None: - state_ids = metadata.state_ids.long() - valid = metadata.valid_state.reshape(-1, - *((1, ) * (state.ndim - 1))) - previous = state.index_select(0, state_ids) - stored = torch.where(valid, value.to(state.dtype), previous) - state.index_copy_(0, state_ids, stored) + ids = torch.where(metadata.valid_state, metadata.state_ids, -1) + _state_scatter(state.unsqueeze(1), ids, torch.zeros_like(ids), value) class CudaKdaImpl(KdaImpl): - """KDA implemented by public FLA convolution/recurrence kernels.""" + """FLA prefill and shared channelwise gated-delta decode.""" def __init__(self): try: @@ -45,7 +43,8 @@ def __init__(self): causal_conv1d_fwd, causal_conv1d_update, ) - from fla.ops.kda import chunk_kda, fused_recurrent_kda + from fla.ops.kda import chunk_kda + from fla.ops.kda.gate import kda_gate_fwd except (ImportError, AttributeError) as exc: raise ImportError( 'GLM-5.3 KDA requires flash-linear-attention==0.5.2.' @@ -53,7 +52,10 @@ def __init__(self): self.causal_conv1d_fwd = causal_conv1d_fwd self.causal_conv1d_update = causal_conv1d_update self.chunk_kda = chunk_kda - self.fused_recurrent_kda = fused_recurrent_kda + from lmdeploy.pytorch.kernels.cuda.gated_delta_rule import fused_recurrent_gated_delta_rule + self.kda_gate = kda_gate_fwd + self.recurrent_func = fused_recurrent_gated_delta_rule + self.fused_recurrent_kda = self._decode_recurrent register_step_metadata_impl(self) def get_step_metadata_provider(self): @@ -62,18 +64,21 @@ def get_step_metadata_provider(self): def _forward_spec(self, mixed_qkv, raw_gate, raw_beta, conv_state, recurrent_state, metadata, **kwargs): - """Reuse single-step FLA kernels, checkpointing every verified token. + """Checkpoint every verified token at its accepted-history ring slot. State is addressed by accepted history length, not the last proposed length. The ring therefore also handles zero/partial acceptance and request reordering without a scheduler-side rollback hook. """ + if metadata.is_decoding: + return self._forward_spec_decode(mixed_qkv, raw_gate, raw_beta, + conv_state, recurrent_state, metadata, **kwargs) history = metadata.cache_seqlens.long() ring_size = metadata.num_spec_tokens + 1 - ids = metadata.state_ids.long() + ids = torch.where(metadata.valid_state, metadata.state_ids, -1).long() read_slot = history.remainder(ring_size) - conv = conv_state[ids, read_slot].clone() - recurrent = recurrent_state[ids, read_slot].clone() + conv = _state_select(conv_state, ids, read_slot) + recurrent = _state_select(recurrent_state, ids, read_slot) local = copy(metadata) local.num_spec_tokens = 0 local.spec_state_offsets = None @@ -82,34 +87,65 @@ def _forward_spec(self, mixed_qkv, raw_gate, raw_beta, conv_state, def store(state, values, lengths): slots = lengths.remainder(ring_size) - previous = state[ids, slots] - valid = metadata.valid_state.reshape(-1, *([1] * (values.ndim - 1))) - state[ids, slots] = torch.where(valid, values.to(state.dtype), previous) - - if not metadata.is_decoding: - output = self.forward(mixed_qkv, raw_gate, raw_beta, - conv_state=conv, recurrent_state=recurrent, - metadata=local, **kwargs) - lengths = history + metadata.cu_seqlens.diff() - store(conv_state, conv, lengths) - store(recurrent_state, recurrent, lengths) - return output - - batch_size = ids.numel() - steps = mixed_qkv.size(1) // batch_size - if steps > ring_size: + _state_scatter(state, ids, slots, values) + + output = self.forward(mixed_qkv, raw_gate, raw_beta, + conv_state=conv, recurrent_state=recurrent, + metadata=local, **kwargs) + lengths = history + metadata.cu_seqlens.diff() + store(conv_state, conv, lengths) + store(recurrent_state, recurrent, lengths) + return output + + def _decode_recurrent(self, q, k, v, g, beta, A_log, dt_bias, initial_state, + output_final_state=True, lower_bound=None, **kwargs): + """Keep AR and MTP on the same recurrence and gate arithmetic.""" + gate = self.kda_gate(g, A_log, dt_bias, lower_bound=lower_bound) + return self.recurrent_func(q, k, v, g=gate, beta=beta.float().sigmoid(), + initial_state=initial_state, output_final_state=output_final_state, + use_qk_l2norm_in_kernel=True, transpose_state_layout=True) + + def _forward_spec_decode(self, mixed_qkv, raw_gate, raw_beta, conv_state, + recurrent_state, metadata, **kwargs): + """Batch convolution windows and verify all tokens in one recurrence. + + The recurrence is parallel across state tiles, not across causally + dependent timesteps. Each timestep is saved for partial acceptance. + """ + ids = metadata.state_ids.long() + batch = ids.numel() + steps = mixed_qkv.size(1) // batch + ring = metadata.num_spec_tokens + 1 + if steps > ring: raise ValueError('KDA verification exceeds the configured state ring.') - inputs = [x.unflatten(1, (batch_size, steps)) - for x in (mixed_qkv, raw_gate, raw_beta)] - outputs = [] - for step in range(steps): - output = self.forward(*(x[:, :, step].contiguous() for x in inputs), - conv_state=conv, recurrent_state=recurrent, - metadata=local, **kwargs) - store(conv_state, conv, history + step + 1) - store(recurrent_state, recurrent, history + step + 1) - outputs.append(output) - return torch.stack(outputs, dim=2).flatten(1, 2) + history = metadata.cache_seqlens + signed_ids = torch.where(metadata.valid_state, ids, -1) + conv = _state_select(conv_state, signed_ids, history.long().remainder(ring)) + values = mixed_qkv.reshape(batch, steps, -1).transpose(1, 2) + width = conv.shape[-1] + windows = torch.cat((conv, values), dim=-1).unfold(-1, width, 1) + windows = windows[:, :, :steps].permute(0, 2, 1, 3).reshape(batch * steps, -1, width).contiguous() + weight = kwargs['conv_weight'] + if weight.ndim == 3: + if weight.size(1) != 1: + raise ValueError('KDA depthwise convolution weight must have shape [D, 1, K].') + weight = weight.squeeze(1) + mixed, conv_out = self.causal_conv1d_update(mixed_qkv, windows, weight=weight, + bias=kwargs['conv_bias'], activation='silu') + slots = (history[:, None] + torch.arange(1, steps + 1, device=ids.device)).remainder(ring) + _state_scatter(conv_state, signed_ids[:, None].expand(-1, steps).reshape(-1).contiguous(), + slots.flatten().contiguous(), conv_out) + heads, dim = kwargs['num_heads'], kwargs['head_dim'] + q, k, v = [x.reshape(batch, steps, heads, dim).contiguous() + for x in mixed.split(heads * dim, dim=-1)] + gate = self.kda_gate(raw_gate.reshape(batch, steps, heads, dim).contiguous(), + kwargs['a_log'], kwargs['dt_bias'], lower_bound=kwargs['lower_bound']) + beta = raw_beta.reshape(batch, steps, heads).float().sigmoid() + output, _ = self.recurrent_func(q, k, v, g=gate, beta=beta, + initial_state=recurrent_state, state_indices=signed_ids, + cache_seqlens=history, output_final_state=True, + use_qk_l2norm_in_kernel=True, transpose_state_layout=True) + return output.reshape(1, batch * steps, heads, dim) def _conv( self, diff --git a/lmdeploy/pytorch/backends/cuda/kpool.py b/lmdeploy/pytorch/backends/cuda/kpool.py index b2556e3c2b..2fd78fab41 100644 --- a/lmdeploy/pytorch/backends/cuda/kpool.py +++ b/lmdeploy/pytorch/backends/cuda/kpool.py @@ -110,6 +110,9 @@ def kpool_select_groups_cuda( fill=-1, descending=True, sorted=False, + # The BF16 sparse-MLA consumer is sensitive to index tile ordering. + # Keep the same selected set/order across AR and MTP verification. + deterministic=True, ) diff --git a/lmdeploy/pytorch/engine/logits_process.py b/lmdeploy/pytorch/engine/logits_process.py index fb0024f992..ed74a94d6b 100644 --- a/lmdeploy/pytorch/engine/logits_process.py +++ b/lmdeploy/pytorch/engine/logits_process.py @@ -498,49 +498,48 @@ async def __call__(self, scores: torch.Tensor) -> torch.Tensor: return scores, logprobs - @torch.inference_mode() - def sampling(self, logits: torch.Tensor): - """sampling.""" + def _filter_sorted_logits(self, logits: torch.Tensor): + """Shared top-k/top-p/min-p policy for sampling and verification.""" sampling_inputs = self.sampling_inputs - - def __random_sampling(scores: torch.Tensor, indices: torch.LongTensor): - """Random sampling.""" - max_topk = sampling_inputs.max_top_k - top_k = sampling_inputs.top_k - if max_topk <= 0: - max_topk = scores.size(1) - if top_k is not None: - top_k = torch.masked_fill(top_k, top_k <= 0, max_topk) - + max_topk = sampling_inputs.max_top_k + top_k = sampling_inputs.top_k + # Sorting is only needed when the full vocabulary can be sampled. + if max_topk <= 0: + scores, indices = logits.sort(1, descending=True) if top_k is not None: - scores = _filter_topk_sorted_(scores, top_k) - - top_p = sampling_inputs.top_p - if top_p is not None: - scores = _filter_topp_sorted_(scores, top_p) - - min_p = sampling_inputs.min_p - if min_p is not None: - scores = _filter_minp_sorted_(scores, min_p) + top_k = torch.masked_fill(top_k, top_k <= 0, scores.size(1)) + else: + scores, indices = _torch_topk(logits, max_topk, dim=1) + if top_k is not None: + scores = _filter_topk_sorted_(scores, top_k) + if sampling_inputs.top_p is not None: + scores = _filter_topp_sorted_(scores, sampling_inputs.top_p) + if sampling_inputs.min_p is not None: + scores = _filter_minp_sorted_(scores, sampling_inputs.min_p) + return scores, indices - softmax_scores = scores.softmax(1) + @torch.inference_mode() + def filter_logits(self, logits: torch.Tensor): + """Apply sampling filters in vocabulary order without modifying logits. - seeds = sampling_inputs.random_seeds - offsets = sampling_inputs.random_offsets - return _multinomial_sampling(softmax_scores, seeds, offsets, indices) + Speculative verification needs the same target distribution as AR, + but its backend consumes logits rather than sampled token IDs. + """ + inputs = self.sampling_inputs + if inputs.max_top_k <= 0 and inputs.top_k is None and inputs.top_p is None and inputs.min_p is None: + return logits + scores, indices = self._filter_sorted_logits(logits) + return torch.full_like(logits, -float('inf')).scatter_(1, indices, scores) + @torch.inference_mode() + def sampling(self, logits: torch.Tensor): + """sampling.""" + sampling_inputs = self.sampling_inputs if sampling_inputs.max_top_k == 1: - result = logits.argmax(-1) - else: - # sort logits is too slow. and we only need topk logits - max_topk = sampling_inputs.max_top_k - if max_topk <= 0: - scores, indices = logits.sort(1, descending=True) - else: - scores, indices = _torch_topk(logits, max_topk, dim=1) - result = __random_sampling(scores, indices) - - return result + return logits.argmax(-1) + scores, indices = self._filter_sorted_logits(logits) + return _multinomial_sampling(scores.softmax(1), sampling_inputs.random_seeds, + sampling_inputs.random_offsets, indices) @torch.inference_mode() def compute_logprobs(self, raw_logprobs: torch.Tensor, token_ids: torch.LongTensor): diff --git a/lmdeploy/pytorch/kernels/cuda/gated_delta_rule.py b/lmdeploy/pytorch/kernels/cuda/gated_delta_rule.py index 2ec3109089..481829bd7d 100644 --- a/lmdeploy/pytorch/kernels/cuda/gated_delta_rule.py +++ b/lmdeploy/pytorch/kernels/cuda/gated_delta_rule.py @@ -223,13 +223,16 @@ def load_value_tile(Value: T.Buffer, v_local: T.Buffer, b_id, seq_id, hv_id, v_o @T.macro def update_recurrent_state(h_local: T.Buffer, k_local: T.Buffer, v_local: T.Buffer, g_exp, beta, k_per_thr: int, - v_per_warp: int) -> None: + v_per_warp: int, channelwise_g: bool = False) -> None: """Apply one gated delta-rule token update to a warp-local state tile.""" for i in T.Unroll(v_per_warp): hk = T.alloc_var(T.float32) hk = 0 for j in T.Unroll(k_per_thr): - h_local[j, i] = h_local[j, i] * g_exp + if channelwise_g: + h_local[j, i] = h_local[j, i] * g_exp[j] + else: + h_local[j, i] = h_local[j, i] * g_exp hk += h_local[j, i] * k_local[j] hk = T.warp_reduce_sum(hk) v_delta = (v_local[i] - hk) * beta @@ -275,7 +278,8 @@ def fused_recurrent_gated_delta_rule_fwd(SEQLEN, use_state_indices: bool = False, is_circular_buffer: bool = False, transpose_state_layout: bool = False, - num_warps: int = 1): + num_warps: int = 1, + channelwise_g: bool = False): """Build the layout-specific recurrent GDR TileLang kernel. Common compile-time metadata is computed once here. The only structural branch is the returned T.prim_func body, @@ -310,7 +314,7 @@ def fused_recurrent_gated_delta_rule_fwd(SEQLEN, num_waves = T.ceildiv(target_v_per_cta, v_per_warp * num_warps) v_per_cta = v_per_warp * num_warps * num_waves use_coalesced_circular_write = is_circular_buffer and num_waves == 1 - use_shared_token_inputs = SEQLEN <= 8 + use_shared_token_inputs = SEQLEN <= 8 and not channelwise_g write_circular_state = output_final_state and is_circular_buffer write_direct_circular_state = write_circular_state and not use_coalesced_circular_write write_coalesced_circular_state = write_circular_state and use_coalesced_circular_write @@ -319,6 +323,7 @@ def fused_recurrent_gated_delta_rule_fwd(SEQLEN, B = T.dynamic('B') N = B if not use_state_indices else T.dynamic('N') + g_shape = [B, SEQLEN, HV, K] if channelwise_g else [B, SEQLEN, HV] if g_dtype is None: g_dtype = dtype @@ -456,7 +461,7 @@ def fused_recurrent_gated_delta_rule_transposed_main( Key: T.StridedTensor([B, SEQLEN, H, K], dtype=dtype, strides=k_stride), Value: T.StridedTensor([B, SEQLEN, HV, V], dtype=dtype, strides=v_stride), Out: T.Tensor([B, SEQLEN, HV, V], dtype=dtype), - G: T.Tensor([B, SEQLEN, HV], dtype=g_dtype), + G: T.Tensor(g_shape, dtype=g_dtype), Beta: T.Tensor([B, SEQLEN, HV], dtype=beta_dtype), State: T.StridedTensor([N, NUM_STATE, HV, V, K], dtype=state_dtype, strides=state_stride), StateIndices: T.Tensor([B], dtype=torch.int64) = None, @@ -522,7 +527,7 @@ def fused_recurrent_gated_delta_rule_transposed_main( g_exp = g_exp_smem[seq_id] beta = beta_smem[seq_id] else: - if use_g: + if use_g and not channelwise_g: if lane_id == 0: g = T.cast(G[b_id, seq_id, hv_id], T.float32) g_exp = T.exp(g) @@ -542,7 +547,14 @@ def fused_recurrent_gated_delta_rule_transposed_main( v_local = T.alloc_local([v_per_warp], dtype) load_value_tile(Value, v_local, b_id, seq_id, hv_id, v_off, V, v_per_warp, data_vw) - update_recurrent_state(h_local, k_local, v_local, g_exp, beta, k_per_thr, v_per_warp) + if channelwise_g: + g_local = T.alloc_local([k_per_thr], T.float32) + for j in T.Unroll(k_per_thr): + g_local[j] = T.exp(T.cast(G[b_id, seq_id, hv_id, k_off + j], T.float32)) + update_recurrent_state(h_local, k_local, v_local, g_local, beta, k_per_thr, + v_per_warp, channelwise_g=True) + else: + update_recurrent_state(h_local, k_local, v_local, g_exp, beta, k_per_thr, v_per_warp) if write_circular_state: store_transposed_state_tile(State, h_local, state_id, state_update_id, hv_id, k_off, v_off, @@ -592,7 +604,8 @@ def fused_recurrent_gated_delta_rule( q: [B, T, H, K] k: [B, T, H, K] v: [B, T, HV, V] - g: [B, T, HV], optional + g: [B, T, HV], optional. The transposed-state path also supports + channelwise log decay [B, T, HV, K] for KDA. beta: [B, T, HV], optional scale: float, optional initial_state: Tensor, optional. Recurrent state with shape @@ -624,6 +637,10 @@ def fused_recurrent_gated_delta_rule( scale = 1 / (q.shape[-1]**0.5) g_dtype = torch.float32 beta_dtype = torch.float32 + channelwise_g = g is not None and g.ndim == 4 + if channelwise_g: + assert transpose_state_layout, 'Channelwise decay requires transposed state layout' + assert g.shape == (*q.shape[:2], HV, K) if g is not None: assert g.is_contiguous() g_dtype = g.dtype @@ -693,6 +710,7 @@ def fused_recurrent_gated_delta_rule( is_circular_buffer=cache_seqlens is not None, transpose_state_layout=transpose_state_layout, num_warps=num_warps, + channelwise_g=channelwise_g, ) kernel(q, k, v, o, g, beta, final_state, state_indices, cache_seqlens) diff --git a/lmdeploy/pytorch/kernels/cuda/sparse_index_topk.py b/lmdeploy/pytorch/kernels/cuda/sparse_index_topk.py index 24a3d674e9..8b0ab07f41 100644 --- a/lmdeploy/pytorch/kernels/cuda/sparse_index_topk.py +++ b/lmdeploy/pytorch/kernels/cuda/sparse_index_topk.py @@ -33,8 +33,10 @@ } -def _ordered_fp32_key(score): +def _ordered_fp32_key(score, canonical_zero: bool = False): """Map fp32 to an integer key whose unsigned order matches fp32 order.""" + if canonical_zero: + score = T.if_then_else(score == 0, T.float32(0), score) bits = T.reinterpret(score, T.uint32) sign_mask = T.cast(2147483648, T.uint32) all_ones = T.cast(4294967295, T.uint32) @@ -52,7 +54,8 @@ def is_sparse_index_topk_supported(k: int) -> bool: @tilelang.jit(pass_configs=_PASS_CONFIGS) def _sparse_index_topk_byte_radix_kernel(top_k: int, fill: int = _FILL, - threads: int = _THREADS): + threads: int = _THREADS, + deterministic: bool = False): num_tokens = T.dynamic('num_tokens') score_width = T.dynamic('score_width') score_stride = T.dynamic('score_stride') @@ -98,7 +101,7 @@ def sparse_index_topk_byte_radix_kernel_( pos = T.alloc_var(T.int32) pos = tidx while pos < seqlen: - key = _ordered_fp32_key(Scores[row, pos]) + key = _ordered_fp32_key(Scores[row, pos], deterministic) if T.bitwise_and(key, prefix_mask) == prefix_key: bin_u32 = T.bitwise_and(key >> shift, T.cast(255, T.uint32)) T.atomic_add(histogram[T.cast(bin_u32, T.int32)], 1) @@ -126,43 +129,75 @@ def sparse_index_topk_byte_radix_kernel_( threshold_key = prefix_key - # Reuse shared_state as output counters after threshold search: - # [0] counts scores above threshold, [1] counts scores equal to it. - if tidx == 0: - shared_state[_STATE_EMIT_GT_COUNT] = 0 - shared_state[_STATE_EMIT_EQ_COUNT] = 0 - T.sync_threads() - - out_pos_buf = T.alloc_local((1,), T.int32) - # First emit all scores strictly greater than the threshold. - pos_emit_gt = T.alloc_var(T.int32) - pos_emit_gt = tidx - while pos_emit_gt < seqlen: - key = _ordered_fp32_key(Scores[row, pos_emit_gt]) - if key > threshold_key: - out_pos_buf[0] = T.atomic_add( - shared_state[_STATE_EMIT_GT_COUNT], 1, return_prev=True) - if out_pos_buf[0] < top_k: - Out[row, out_pos_buf[0]] = pos_emit_gt - pos_emit_gt += threads - - T.sync_threads() - - gt_count = shared_state[_STATE_EMIT_GT_COUNT] - - # Then fill remaining slots from scores equal to the threshold. - # Tie order is intentionally unspecified; sparse attention needs - # a valid top-k set, not score-sorted ids. - pos_emit_eq = T.alloc_var(T.int32) - pos_emit_eq = tidx - while pos_emit_eq < seqlen: - key = _ordered_fp32_key(Scores[row, pos_emit_eq]) - if key == threshold_key: - out_pos_buf[0] = gt_count + T.atomic_add( - shared_state[_STATE_EMIT_EQ_COUNT], 1, return_prev=True) - if out_pos_buf[0] < top_k: - Out[row, out_pos_buf[0]] = pos_emit_eq - pos_emit_eq += threads + if deterministic: + # Preserve the radix threshold but emit in logical-index + # order. Atomic output offsets vary between launches and + # change sparse attention's BF16 softmax tile boundaries. + # ``rank`` is the number of cutoff ties still needed; take + # the lowest logical ids, independent of batch/CTA order. + equal_scan = T.alloc_shared((threads,), T.int32) + selected_scan = T.alloc_shared((threads,), T.int32) + emitted = T.alloc_var(T.int32) + equal_seen = T.alloc_var(T.int32) + emitted = 0 + equal_seen = 0 + for tile in T.serial(T.ceildiv(seqlen, threads)): + position = tile * threads + tidx + stable_key = T.alloc_var(T.uint32) + stable_key = T.cast(0, T.uint32) + if position < seqlen: + stable_key = _ordered_fp32_key(Scores[row, position], True) + equal_scan[tidx] = T.cast(position < seqlen and stable_key == threshold_key, T.int32) + T.sync_threads() + T.cumsum(equal_scan) + T.sync_threads() + selected = position < seqlen and ( + stable_key > threshold_key or + (stable_key == threshold_key and equal_seen + equal_scan[tidx] <= rank)) + selected_scan[tidx] = T.cast(selected, T.int32) + T.sync_threads() + T.cumsum(selected_scan) + T.sync_threads() + if selected: + Out[row, emitted + selected_scan[tidx] - 1] = position + emitted += selected_scan[threads - 1] + equal_seen += equal_scan[threads - 1] + T.sync_threads() + + else: + # Preserve the existing unordered path for other callers. + # Reuse shared_state as counters for above/equal cutoff. + if tidx == 0: + shared_state[_STATE_EMIT_GT_COUNT] = 0 + shared_state[_STATE_EMIT_EQ_COUNT] = 0 + T.sync_threads() + + out_pos_buf = T.alloc_local((1,), T.int32) + pos_emit_gt = T.alloc_var(T.int32) + pos_emit_gt = tidx + while pos_emit_gt < seqlen: + key = _ordered_fp32_key(Scores[row, pos_emit_gt]) + if key > threshold_key: + out_pos_buf[0] = T.atomic_add( + shared_state[_STATE_EMIT_GT_COUNT], 1, return_prev=True) + if out_pos_buf[0] < top_k: + Out[row, out_pos_buf[0]] = pos_emit_gt + pos_emit_gt += threads + + T.sync_threads() + gt_count = shared_state[_STATE_EMIT_GT_COUNT] + + # Tie order remains unspecified in the default path. + pos_emit_eq = T.alloc_var(T.int32) + pos_emit_eq = tidx + while pos_emit_eq < seqlen: + key = _ordered_fp32_key(Scores[row, pos_emit_eq]) + if key == threshold_key: + out_pos_buf[0] = gt_count + T.atomic_add( + shared_state[_STATE_EMIT_EQ_COUNT], 1, return_prev=True) + if out_pos_buf[0] < top_k: + Out[row, out_pos_buf[0]] = pos_emit_eq + pos_emit_eq += threads return sparse_index_topk_byte_radix_kernel_ @@ -173,12 +208,15 @@ def sparse_index_topk(scores: torch.Tensor, k: int, fill: int = _FILL, descending: bool = True, - sorted: bool = False) -> torch.Tensor: + sorted: bool = False, + deterministic: bool = False) -> torch.Tensor: """Return top-k score indices for padded sparse-index score rows. The returned ids are packed but not score-sorted. Sparse attention consumes them as a set of valid KV positions; avoiding final sorting is the point of this selector. Rows shorter than ``k`` are padded with ``fill``. + With ``deterministic=True``, cutoff ties prefer lower logical ids and the + selected ids are emitted in ascending index order (not score order). """ if not descending: raise ValueError('sparse_index_topk only supports descending=True.') @@ -199,5 +237,5 @@ def sparse_index_topk(scores: torch.Tensor, kv_seqlens = kv_seqlens.contiguous() out = torch.empty((num_tokens, k), device=scores.device, dtype=torch.int32) - _sparse_index_topk_byte_radix_kernel(k, fill, _THREADS)(scores, kv_seqlens, out) + _sparse_index_topk_byte_radix_kernel(k, fill, _THREADS, deterministic)(scores, kv_seqlens, out) return out diff --git a/lmdeploy/pytorch/models/glm5_next.py b/lmdeploy/pytorch/models/glm5_next.py index 886e50850a..5a184a28dc 100644 --- a/lmdeploy/pytorch/models/glm5_next.py +++ b/lmdeploy/pytorch/models/glm5_next.py @@ -70,6 +70,7 @@ Glm4vVisionPatchEmbed, Glm4vVisionRotaryEmbedding, ) +from .glm_moe_dsa import DSATopKIndicesBuffer from .glm_moe_dsa_mtp import GlmMoeDsaMTPModel, GlmMoeDsaMultiTokenPredictor from .qwen3_vl import Qwen3VLInputProcessor from .utils.model import build_embedding, vlm_model @@ -435,6 +436,8 @@ class Glm5NextMLP(DeepseekV2MLP): def __init__(self, config: Any, *args, **kwargs): super().__init__(config, *args, **kwargs) self.swiglu_limit = config.swiglu_limit + if get_dist_manager().current_config().dp == 1: + self.down_proj.tp_reduce_dtype = torch.float32 def forward(self, x: torch.Tensor) -> torch.Tensor: gate_up = self.gate_up_proj(x) @@ -518,6 +521,11 @@ class Glm5NextMoE(DeepseekV2MoE): def __init__(self, config: Any, layer_idx: int, *args, **kwargs): kwargs.setdefault('prefix', f'model.layers.{layer_idx}.mlp') super().__init__(config, layer_idx, *args, **kwargs) + # Keep the shared+routed local sum and the generic expert kernels. + # Promote only the final TP collective: BF16 collective reduction + # order depends on message size (AR versus multi-token verification). + self._fp32_tp_reduce = self._all_reduce + self._all_reduce = False if self.gate.fake_eplb or self.gate.eplb_dispatch_info is not None: raise RuntimeError( 'The GLM-5.3 router does not permit fake ' @@ -532,7 +540,13 @@ def forward( if all_routed_experts is not None: raise RuntimeError( 'GLM-5.3 routed-expert capture is not supported.') - return super().forward(hidden_states, all_routed_experts=None) + out = super().forward(hidden_states, all_routed_experts=None) + if self._fp32_tp_reduce: + output_dtype = out.dtype + out = out.float() + dist.all_reduce(out, group=self.experts.tp_group) + out = out.to(output_dtype) + return out def _load_vector_shard(param: nn.Parameter, @@ -694,6 +708,8 @@ def __init__(self, is_tp=True, all_reduce=all_reduce, ) + if get_dist_manager().current_config().dp == 1: + self.o_proj.tp_reduce_dtype = torch.float32 self.kda = Kda() def forward(self, hidden_states: torch.Tensor, @@ -746,6 +762,8 @@ def __init__(self, device=device, all_reduce=all_reduce, prefix=f'model.layers.{layer_idx}.self_attn') + if get_dist_manager().current_config().dp == 1: + self.o_proj.tp_reduce_dtype = torch.float32 # DeepSeek keeps these latent-norm parameters in FP32. GLM-5.3's # checkpoint and SGLang runtime keep them in the activation dtype; # rebuild only these two containers before weight loading. @@ -769,7 +787,8 @@ def __init__(self, # Keep the checkpoint's BF16 KV-B projection alongside the absorbed # KC/VC views. Short prefill uses the former to reproduce SGLang's - # decompressed dense MHA; decode continues to use KC/VC. + # decompressed dense MHA; decode continues to use KC/VC. Both must + # shard by attention TP, including when attention DP is enabled. self.kv_b_proj = build_colwise_linear( self.kv_lora_rank, self.num_heads * (self.qk_nope_head_dim + self.v_head_dim), @@ -778,7 +797,6 @@ def __init__(self, device=device, is_tp=True, quant_config=None, - dp_disable_tp=True, ) self.prefill_attn_fwd = FlashAttention( self.num_heads, @@ -1134,13 +1152,23 @@ def _kpool_indices( state_ids: torch.Tensor, attn_metadata: Any, return_indices: bool, + topk_indices_buffer: DSATopKIndicesBuffer | None = None, + skip_topk: bool = False, ) -> torch.Tensor | None: indexer_k_cache = self._update_kpool_cache( hidden_states, tail_state, state_ids, attn_metadata) - if not return_indices: + if topk_indices_buffer is not None and skip_topk: + return (topk_indices_buffer.read(hidden_states.size(1), hidden_states.device) + if return_indices else None) + if not return_indices and topk_indices_buffer is None: return None - return self._select_kpool_indices( + indices = self._select_kpool_indices( hidden_states, q_lora, indexer_k_cache, attn_metadata) + if topk_indices_buffer is not None: + # MTP needs seed indices even for dense short prefill: subsequent + # draft steps reuse its last-token rows through the shared proposer. + indices = topk_indices_buffer.write(indices) + return indices if return_indices else None def _absorbed_query(self, query: torch.Tensor, num_heads: int) -> torch.Tensor: @@ -1199,6 +1227,8 @@ def forward( attn_metadata: Any = None, kpool_tail_state: Sequence[torch.Tensor] | None = None, state_ids: torch.Tensor | None = None, + topk_indices_buffer: DSATopKIndicesBuffer | None = None, + skip_topk: bool = False, ) -> torch.Tensor: dist_ctx = get_dist_manager().current_context() num_heads = self.num_heads // dist_ctx.dist_config.attn_tp @@ -1221,6 +1251,8 @@ def forward( state_ids, attn_metadata, return_indices=use_sparse, + topk_indices_buffer=topk_indices_buffer, + skip_topk=skip_topk, ) if not use_sparse: return self._forward_prefill_mha( @@ -1258,6 +1290,8 @@ def forward( state_ids, attn_metadata, return_indices=True, + topk_indices_buffer=topk_indices_buffer, + skip_topk=skip_topk, ) # GLM has no RoPE tail, so the absorbed query contains exactly 512 # values; the cache retains its 576-wide FlashMLA storage alignment. @@ -1329,7 +1363,8 @@ def __init__(self, device=device) self.hc_prepost = HcPrePost(config.hc_mult, config.hc_sinkhorn_iters, - config.hc_eps) + config.hc_eps, + avoid_gemv=True) mix_hc = (2 + config.hc_mult) * config.hc_mult hc_dim = config.hc_mult * config.hidden_size self.hc_attn_fn = nn.Parameter(torch.empty(mix_hc, @@ -1956,10 +1991,13 @@ def __init__(self, config, layer_idx, dtype=None, device=None): dtype=dtype, device=device) def forward(self, hidden_states, rotary_pos_emb, past_key_value, - attn_metadata=None, **kwargs): + attn_metadata=None, topk_indices_buffer=None, + skip_topk=False, **kwargs): residual = hidden_states hidden_states = self.input_layernorm(hidden_states) - hidden_states = self.self_attn(hidden_states, past_key_value, attn_metadata) + hidden_states = self.self_attn( + hidden_states, past_key_value, attn_metadata, + topk_indices_buffer=topk_indices_buffer, skip_topk=skip_topk) hidden_states, residual = self.post_attention_layernorm(hidden_states, residual) return self.mlp(hidden_states), residual @@ -1980,8 +2018,10 @@ def __init__(self, config, ctx_mgr, dtype=None, device=None): self.model = GlmMoeDsaMultiTokenPredictor( self.config, dtype=dtype, device=device, decoder_layer_cls=Glm5NextMTPDecoderLayer) - self.uses_dsa_topk_buffer = False - self.topk_indices_buffer = None + self.uses_dsa_topk_buffer = getattr(self.config, 'index_share_for_mtp_iteration', False) + self.topk_indices_buffer = ( + DSATopKIndicesBuffer(self.config.index_topk + self.config.index_kpool - 1) + if self.uses_dsa_topk_buffer else None) self._load_buffers = {} def prepare_inputs_for_generation(self, past_key_values, inputs_embeds=None, diff --git a/lmdeploy/pytorch/models/glm_moe_dsa_mtp.py b/lmdeploy/pytorch/models/glm_moe_dsa_mtp.py index fd5ed7dd77..7fc18402b2 100644 --- a/lmdeploy/pytorch/models/glm_moe_dsa_mtp.py +++ b/lmdeploy/pytorch/models/glm_moe_dsa_mtp.py @@ -239,8 +239,14 @@ def forward( ) def get_cudagraph_extra_key(self, skip_topk: bool = False, **kwargs) -> tuple: - """Separate graphs that compute and reuse DSA top-k indices.""" - return (skip_topk, ) + """Separate seed/reuse graphs and invalidate captured grown buffers.""" + buffer = getattr(self, 'topk_indices_buffer', None) + # The final shifted MTP chunk can exceed max_prefill_token_num by one; + # large multimodal spans can grow it further. A captured graph retains + # the old allocation, so do not replay it after the buffer grows. + capacity = (0 if buffer is None or buffer.indices is None + else buffer.indices.size(0)) + return (capacity, skip_topk) def prepare_inputs_for_generation( self, diff --git a/lmdeploy/pytorch/nn/hc_prepost.py b/lmdeploy/pytorch/nn/hc_prepost.py index 97f0c717ab..3106b9cb0a 100644 --- a/lmdeploy/pytorch/nn/hc_prepost.py +++ b/lmdeploy/pytorch/nn/hc_prepost.py @@ -11,8 +11,9 @@ class HcPrePost(nn.Module): """DeepSeek-V4 hyper-connection pre/post reduction wrapper.""" - def __init__(self, hc_mult: int, sinkhorn_iters: int = 20, eps: float = 1e-6): + def __init__(self, hc_mult: int, sinkhorn_iters: int = 20, eps: float = 1e-6, *, avoid_gemv: bool = False): super().__init__() + self.avoid_gemv = avoid_gemv self.impl = get_backend().build_op( HCPrePostBuildSpec(hc_mult=hc_mult, sinkhorn_iters=sinkhorn_iters, eps=eps), enable_deterministic=get_build_model_context().enable_deterministic, @@ -29,7 +30,13 @@ def pre( from lmdeploy.pytorch.nn.norm import rms_scale shape, dtype = x.size(), x.dtype x = x.flatten(2).float() - mixes = rms_scale(F.linear(x, hc_fn), x, eps=norm_eps) + if self.avoid_gemv and x.size(0) == 1 and x.size(1) == 1: + # Single-token decode otherwise selects GEMV, whose reduction + # order can differ from multi-token speculative verification. + mixes = F.linear(F.pad(x, (0, 0, 0, 1)), hc_fn)[:, :1].contiguous() + else: + mixes = F.linear(x, hc_fn) + mixes = rms_scale(mixes, x, eps=norm_eps) return self.impl.pre(x.view(shape), mixes, hc_scale, hc_base, dtype) def pre_reduce(self, x: torch.Tensor, pre: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: diff --git a/lmdeploy/pytorch/nn/linear/base.py b/lmdeploy/pytorch/nn/linear/base.py index 03ec93bb0f..77c166ef3e 100644 --- a/lmdeploy/pytorch/nn/linear/base.py +++ b/lmdeploy/pytorch/nn/linear/base.py @@ -111,6 +111,10 @@ def __slice_and_gather(): class LinearBase(nn.Module): """Base class for linear layers.""" + # Optional accumulation dtype for the unfused TP/LoRA reduction. Keep + # backend-fused communication and DP_TP unchanged by default. + tp_reduce_dtype = None + def __init__( self, dtype: torch.dtype | None = None, @@ -198,7 +202,7 @@ def _forward_default(self, x, all_reduce: bool, tp_sizes: list[int]): raise NotImplementedError('This method should be implemented in subclasses.') def _forward_lora(self, x, tp_sizes: list[int] = None): - """Forward with LoRA.""" + """Local projection and optional LoRA, followed by TP reduction.""" out = self._forward_default(x, False, tp_sizes) for lora_adapter in self.lora_adapters.values(): @@ -207,7 +211,11 @@ def _forward_lora(self, x, tp_sizes: list[int] = None): if self.tp_mode == TPMode.DP_TP: out = reduce_scatter_by_tp_sizes(out, self.tp_rank, tp_sizes, group=self.tp_group) else: + output_dtype = out.dtype + if self.tp_reduce_dtype is not None: + out = out.to(self.tp_reduce_dtype) dist.all_reduce(out, group=self.tp_group) + out = out.to(output_dtype) return out def _forward_dp_tp(self, x): @@ -232,7 +240,7 @@ def forward(self, x): if self.tp > 1 and self.tp_mode == TPMode.DP_TP: return self._forward_dp_tp(x) - if len(self.lora_adapters) == 0: + if len(self.lora_adapters) == 0 and not (self.all_reduce and self.tp_reduce_dtype is not None): return self._forward_default(x, self.all_reduce, None) else: return self._forward_lora(x) diff --git a/lmdeploy/pytorch/spec_decode/reject_sampler.py b/lmdeploy/pytorch/spec_decode/reject_sampler.py index 18e5100cbe..092c0052be 100644 --- a/lmdeploy/pytorch/spec_decode/reject_sampler.py +++ b/lmdeploy/pytorch/spec_decode/reject_sampler.py @@ -46,7 +46,11 @@ def forward( bonus_logits = target_logits[:, -1] bonus_token_ids = FusedLogitsProcessor( bonus_sampling_inputs).sampling(bonus_logits) - target_draft_logits = target_logits[:, :-1].contiguous() + # Verification/recovery must use the AR sampling distribution too. + # Filter out of place so bonus sampling and returned logits are intact. + filtered_logits = FusedLogitsProcessor( + expanded_sampling_inputs).filter_logits(target_logits.flatten(0, 1)) + target_draft_logits = filtered_logits.view_as(target_logits)[:, :-1].contiguous() is_greedy = None if bonus_sampling_inputs.has_greedy: From 41563c7c639e172da775f2e73dd253977f443a40 Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Tue, 22 Sep 2026 09:56:46 +0000 Subject: [PATCH 06/39] refactor(pytorch): reuse causal convolution cache for GLM MTP --- lmdeploy/pytorch/backends/cuda/kda.py | 32 +++++++++++-------- lmdeploy/pytorch/configurations/glm5_next.py | 6 ++-- .../pytorch/kernels/cuda/causal_conv1d.py | 12 ++++--- 3 files changed, 30 insertions(+), 20 deletions(-) diff --git a/lmdeploy/pytorch/backends/cuda/kda.py b/lmdeploy/pytorch/backends/cuda/kda.py index aaef5a2180..ff3511cec1 100644 --- a/lmdeploy/pytorch/backends/cuda/kda.py +++ b/lmdeploy/pytorch/backends/cuda/kda.py @@ -52,7 +52,9 @@ def __init__(self): self.causal_conv1d_fwd = causal_conv1d_fwd self.causal_conv1d_update = causal_conv1d_update self.chunk_kda = chunk_kda + from lmdeploy.pytorch.kernels.cuda.causal_conv1d import causal_conv1d_update as shared_conv_update from lmdeploy.pytorch.kernels.cuda.gated_delta_rule import fused_recurrent_gated_delta_rule + self.shared_conv_update = shared_conv_update self.kda_gate = kda_gate_fwd self.recurrent_func = fused_recurrent_gated_delta_rule self.fused_recurrent_kda = self._decode_recurrent @@ -77,7 +79,13 @@ def _forward_spec(self, mixed_qkv, raw_gate, raw_beta, conv_state, ring_size = metadata.num_spec_tokens + 1 ids = torch.where(metadata.valid_state, metadata.state_ids, -1).long() read_slot = history.remainder(ring_size) - conv = _state_select(conv_state, ids, read_slot) + # FLA prefill consumes a chronological window; the persistent cache + # uses the same compact token ring as Qwen3.5's causal convolution. + conv_cache = _select_state(conv_state, metadata) + width = kwargs['conv_weight'].shape[-1] + offsets = torch.arange(-width, 0, device=ids.device) + read_offsets = (history[:, None] + offsets).remainder(conv_state.size(-1)) + conv = conv_cache.gather(2, read_offsets[:, None].expand(-1, conv_cache.size(1), -1)) recurrent = _state_select(recurrent_state, ids, read_slot) local = copy(metadata) local.num_spec_tokens = 0 @@ -93,7 +101,9 @@ def store(state, values, lengths): conv_state=conv, recurrent_state=recurrent, metadata=local, **kwargs) lengths = history + metadata.cu_seqlens.diff() - store(conv_state, conv, lengths) + write_offsets = (lengths[:, None] + offsets).remainder(conv_state.size(-1)) + conv_cache.scatter_(2, write_offsets[:, None].expand(-1, conv_cache.size(1), -1), conv) + _store_state(conv_state, conv_cache, metadata) store(recurrent_state, recurrent, lengths) return output @@ -107,7 +117,7 @@ def _decode_recurrent(self, q, k, v, g, beta, A_log, dt_bias, initial_state, def _forward_spec_decode(self, mixed_qkv, raw_gate, raw_beta, conv_state, recurrent_state, metadata, **kwargs): - """Batch convolution windows and verify all tokens in one recurrence. + """Reuse causal-convolution token rings and one verification recurrence. The recurrence is parallel across state tiles, not across causally dependent timesteps. Each timestep is saved for partial acceptance. @@ -120,21 +130,17 @@ def _forward_spec_decode(self, mixed_qkv, raw_gate, raw_beta, conv_state, raise ValueError('KDA verification exceeds the configured state ring.') history = metadata.cache_seqlens signed_ids = torch.where(metadata.valid_state, ids, -1) - conv = _state_select(conv_state, signed_ids, history.long().remainder(ring)) - values = mixed_qkv.reshape(batch, steps, -1).transpose(1, 2) - width = conv.shape[-1] - windows = torch.cat((conv, values), dim=-1).unfold(-1, width, 1) - windows = windows[:, :, :steps].permute(0, 2, 1, 3).reshape(batch * steps, -1, width).contiguous() + values = mixed_qkv.reshape(batch, steps, -1).transpose(1, 2).contiguous() weight = kwargs['conv_weight'] if weight.ndim == 3: if weight.size(1) != 1: raise ValueError('KDA depthwise convolution weight must have shape [D, 1, K].') weight = weight.squeeze(1) - mixed, conv_out = self.causal_conv1d_update(mixed_qkv, windows, weight=weight, - bias=kwargs['conv_bias'], activation='silu') - slots = (history[:, None] + torch.arange(1, steps + 1, device=ids.device)).remainder(ring) - _state_scatter(conv_state, signed_ids[:, None].expand(-1, steps).reshape(-1).contiguous(), - slots.flatten().contiguous(), conv_out) + mixed = self.shared_conv_update(values, conv_state, weight, + bias=kwargs['conv_bias'], activation='silu', + conv_state_indices=signed_ids.to(torch.int32), + cache_seqlens=history) + mixed = mixed.transpose(1, 2) heads, dim = kwargs['num_heads'], kwargs['head_dim'] q, k, v = [x.reshape(batch, steps, heads, dim).contiguous() for x in mixed.split(heads * dim, dim=-1)] diff --git a/lmdeploy/pytorch/configurations/glm5_next.py b/lmdeploy/pytorch/configurations/glm5_next.py index 2ef73daf4a..fda90cf426 100644 --- a/lmdeploy/pytorch/configurations/glm5_next.py +++ b/lmdeploy/pytorch/configurations/glm5_next.py @@ -244,13 +244,13 @@ def build(cls, hf_config, model_path: str | None = None, **kwargs): # unset also preserves the BF16 latent MLA cache policy. config.mla_index_topk = None config.k_head_dim = text_config.kv_lora_rank + 64 - # Keep a complete state after each verified token. Accepted sequence - # lengths select the correct ring slot after rejection sampling. + # Reuse Qwen3.5's token ring for convolution; recurrent/KPool states + # keep a complete checkpoint after each verified token. ring_shape = (num_spec_tokens + 1,) if num_spec_tokens else () config.state_cache_specs = [ StateCacheSpec( GLM5_KDA_CONV_STATE, - (num_linear_layers, *ring_shape, conv_dim, conv_kernel_size), + (num_linear_layers, conv_dim, conv_kernel_size + num_spec_tokens), torch.bfloat16, ), StateCacheSpec( diff --git a/lmdeploy/pytorch/kernels/cuda/causal_conv1d.py b/lmdeploy/pytorch/kernels/cuda/causal_conv1d.py index b634310918..0dbb7e8e46 100644 --- a/lmdeploy/pytorch/kernels/cuda/causal_conv1d.py +++ b/lmdeploy/pytorch/kernels/cuda/causal_conv1d.py @@ -217,13 +217,15 @@ def causal_conv1d_fn( }, ) def causal_conv1d_update_fwd(hidden_size: int, seqlen: int, state_len: int, width: int, has_bias: bool, activation: str | None, dtype, conv_stride: tuple[int, int, int], is_circular_buffer: bool, - has_state_indices: bool, num_warps: int): + has_state_indices: bool, num_warps: int, weight_dtype=None, bias_dtype=None): """TileLang kernel for causal convolution forward pass. Each thread processes one output position for all channels sequentially. """ num_threads = num_warps * 32 silu_activation = activation in ['silu', 'swish'] + weight_dtype = dtype if weight_dtype is None else weight_dtype + bias_dtype = dtype if bias_dtype is None else bias_dtype advance_len = seqlen batch = T.dynamic('batch') @@ -237,8 +239,8 @@ def causal_conv1d_update_main( Conv_State: T.StridedTensor((conv_batch, hidden_size, state_len), dtype=dtype, strides=(conv_batch_stride, conv_stride[1], conv_stride[2])), - W: T.Tensor((hidden_size, width), dtype=dtype), - Bias: T.Tensor((hidden_size, ), dtype=dtype) = None, + W: T.Tensor((hidden_size, width), dtype=weight_dtype), + Bias: T.Tensor((hidden_size, ), dtype=bias_dtype) = None, Out: T.Tensor((batch, hidden_size, seqlen), dtype=dtype) = None, Cache_seqlens: T.Tensor((batch, ), dtype=T.int32) = None, Conv_state_indices: T.Tensor((batch, ), dtype=T.int32) = None, @@ -368,7 +370,9 @@ def causal_conv1d_update(x, conv_stride=conv_state.stride(), is_circular_buffer=cache_seqlens is not None, has_state_indices=conv_state_indices is not None, - num_warps=num_warps) + num_warps=num_warps, + weight_dtype=weight.dtype, + bias_dtype=bias.dtype if bias is not None else x.dtype) kernel(x, conv_state, weight, bias, out, cache_seqlens, conv_state_indices) From 9642332157d45af32a49196888b963bee60f8f0f Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Tue, 22 Sep 2026 11:05:33 +0000 Subject: [PATCH 07/39] fix(pytorch): honor MoE reduction options across backends --- lmdeploy/pytorch/backends/cuda/moe/default.py | 23 ++++++- lmdeploy/pytorch/backends/moe.py | 2 + .../pytorch/kernels/cuda/moe/fused_moe.py | 9 ++- lmdeploy/pytorch/nn/moe/__init__.py | 8 +++ lmdeploy/pytorch/nn/moe/default.py | 6 +- tests/pytorch/nn/test_moe_options.py | 68 +++++++++++++++++++ 6 files changed, 110 insertions(+), 6 deletions(-) create mode 100644 tests/pytorch/nn/test_moe_options.py diff --git a/lmdeploy/pytorch/backends/cuda/moe/default.py b/lmdeploy/pytorch/backends/cuda/moe/default.py index d0b1c31033..00af6d41d3 100644 --- a/lmdeploy/pytorch/backends/cuda/moe/default.py +++ b/lmdeploy/pytorch/backends/cuda/moe/default.py @@ -22,10 +22,17 @@ class TritonFusedMoEImpl(FusedMoEImpl): """Triton fused moe implementation.""" - def __init__(self, top_k: int, num_experts: int, renormalize: bool = False): + def __init__(self, + top_k: int, + num_experts: int, + renormalize: bool = False, + fp32_acc: bool = False, + output_scale: float = 1.0): self.num_experts = num_experts self.top_k = top_k self.renormalize = renormalize + self.fp32_acc = fp32_acc + self.output_scale = output_scale def update_weights(self, gate_up_weights: torch.Tensor, down_weights: torch.Tensor): gate_up_weights = gate_up_weights.transpose(1, 2).contiguous().transpose(1, 2) @@ -67,7 +74,9 @@ def forward(self, expert_offset=expert_offset, num_experts=num_experts, renormalize=self.renormalize, - act_func=act_func) + act_func=act_func, + fp32_acc=self.fp32_acc, + output_scale=self.output_scale) # modify from dlblas: https://github.com/DeepLink-org/DLBlas @@ -375,11 +384,15 @@ def __init__( num_experts: int, hidden_dim: int, renormalize: bool = False, + fp32_acc: bool = False, + output_scale: float = 1.0, layer_idx: int = 0, out_dtype: torch.dtype = torch.bfloat16, num_max_dispatch_tokens_per_rank: int = 128, ): - super().__init__(top_k, num_experts, renormalize) + super().__init__(top_k, num_experts, renormalize, fp32_acc, output_scale) + if fp32_acc or output_scale != 1.0: + raise NotImplementedError('DeepEP MoE does not support fp32_acc or output_scale.') self.num_experts = num_experts self.ep_size = ep_size self.ep_group = ep_group @@ -540,6 +553,8 @@ def _build_fused_moe(spec: FusedMoEBuildSpec) -> FusedMoEImpl: num_experts=spec.num_experts, hidden_dim=spec.hidden_dim, renormalize=spec.renormalize, + fp32_acc=spec.fp32_acc, + output_scale=spec.output_scale, layer_idx=spec.layer_idx, out_dtype=spec.output_dtype, num_max_dispatch_tokens_per_rank=spec.num_max_dispatch_tokens_per_rank, @@ -548,4 +563,6 @@ def _build_fused_moe(spec: FusedMoEBuildSpec) -> FusedMoEImpl: top_k=spec.top_k, num_experts=spec.num_experts, renormalize=spec.renormalize, + fp32_acc=spec.fp32_acc, + output_scale=spec.output_scale, ) diff --git a/lmdeploy/pytorch/backends/moe.py b/lmdeploy/pytorch/backends/moe.py index 0af6af9dd8..9872aada3d 100644 --- a/lmdeploy/pytorch/backends/moe.py +++ b/lmdeploy/pytorch/backends/moe.py @@ -73,6 +73,8 @@ class FusedMoEBuildSpec(BuildSpec[FusedMoEImpl]): layer_idx: int output_dtype: torch.dtype num_max_dispatch_tokens_per_rank: int + fp32_acc: bool = False + output_scale: float = 1.0 class FusedMoEW8A8Impl(ABC): diff --git a/lmdeploy/pytorch/kernels/cuda/moe/fused_moe.py b/lmdeploy/pytorch/kernels/cuda/moe/fused_moe.py index ca1d19ed8a..ea11c12968 100644 --- a/lmdeploy/pytorch/kernels/cuda/moe/fused_moe.py +++ b/lmdeploy/pytorch/kernels/cuda/moe/fused_moe.py @@ -1030,7 +1030,9 @@ def fused_moe(hidden_states: torch.Tensor, expert_offset: int = 0, num_experts: int = None, renormalize: bool = False, - act_func: Callable = None) -> torch.Tensor: + act_func: Callable = None, + fp32_acc: bool = False, + output_scale: float = 1.0) -> torch.Tensor: """Fused moe.""" M = hidden_states.size(0) E, N, _ = w1.shape @@ -1138,5 +1140,8 @@ def fused_moe(hidden_states: torch.Tensor, reindex_c=True, ) - ret = moe_reduce(intermediate_cache2, topk_weights) + ret = moe_reduce(intermediate_cache2, + topk_weights, + fp32_acc=fp32_acc, + output_scale=output_scale) return ret diff --git a/lmdeploy/pytorch/nn/moe/__init__.py b/lmdeploy/pytorch/nn/moe/__init__.py index c834f5e32c..3b30e6a7e2 100644 --- a/lmdeploy/pytorch/nn/moe/__init__.py +++ b/lmdeploy/pytorch/nn/moe/__init__.py @@ -48,9 +48,13 @@ def build_fused_moe( all_reduce=all_reduce, layer_idx=layer_idx, act_func=act_func, + fp32_acc=fp32_acc, + output_scale=output_scale, ) if quant_method == 'smooth_quant': + if fp32_acc or output_scale != 1.0: + raise NotImplementedError('W8A8 MoE does not support fp32_acc or output_scale.') assert not bias, 'Quant model does not support bias for now.' assert act_func is None, ('Quant model does not support activation function for now.') from .w8a8 import FusedMoEW8A8 @@ -72,6 +76,8 @@ def build_fused_moe( ) if is_static_per_tensor: + if fp32_acc or output_scale != 1.0: + raise NotImplementedError('Static FP8 MoE does not support fp32_acc or output_scale.') assert not bias, ( 'Static FP8 MoE does not support bias.' ) @@ -114,6 +120,8 @@ def build_fused_moe( output_scale=output_scale, ) elif quant_method == 'compressed-tensors': + if fp32_acc or output_scale != 1.0: + raise NotImplementedError('W4A16 MoE does not support fp32_acc or output_scale.') if bias: raise RuntimeError('Compressed-tensors W4A16 routed experts do not support bias.') if act_func is not None: diff --git a/lmdeploy/pytorch/nn/moe/default.py b/lmdeploy/pytorch/nn/moe/default.py index 7705a0f082..4906d967cc 100644 --- a/lmdeploy/pytorch/nn/moe/default.py +++ b/lmdeploy/pytorch/nn/moe/default.py @@ -130,7 +130,9 @@ def __init__(self, device: torch.device | None = None, all_reduce: bool = True, layer_idx: int = 0, - act_func: Callable = None): + act_func: Callable = None, + fp32_acc: bool = False, + output_scale: float = 1.0): device = device or torch.device('cpu') dtype = dtype or torch.float16 @@ -158,6 +160,8 @@ def __init__(self, layer_idx=layer_idx, output_dtype=torch.bfloat16, num_max_dispatch_tokens_per_rank=build_ctx.deep_ep_max_tokens_per_rank, + fp32_acc=fp32_acc, + output_scale=output_scale, ), enable_deterministic=build_ctx.enable_deterministic, ) diff --git a/tests/pytorch/nn/test_moe_options.py b/tests/pytorch/nn/test_moe_options.py new file mode 100644 index 0000000000..7f88b6887a --- /dev/null +++ b/tests/pytorch/nn/test_moe_options.py @@ -0,0 +1,68 @@ +# Copyright (c) OpenMMLab. All rights reserved. +import pytest +import torch + + +def test_build_fused_moe_propagates_reduction_options(monkeypatch): + from lmdeploy.pytorch.nn import moe + + captured = {} + + class FakeFusedMoE: + + def __init__(self, **kwargs): + captured.update(kwargs) + + import lmdeploy.pytorch.nn.moe.default as default_moe + monkeypatch.setattr(default_moe, 'FusedMoE', FakeFusedMoE) + + moe.build_fused_moe(16, + 32, + 4, + 2, + quant_config=None, + fp32_acc=True, + output_scale=2.5) + + assert captured['fp32_acc'] is True + assert captured['output_scale'] == 2.5 + + +@pytest.mark.parametrize('quant_method', ['smooth_quant', 'fp8', 'compressed-tensors']) +def test_build_fused_moe_rejects_unsupported_reduction_options(monkeypatch, quant_method): + from lmdeploy.pytorch.nn import moe + + class FakeQuantConfig: + quant_dtype = None + activation_scheme = 'static' + weight_block_size = None + bits = 4 + group_size = 128 + + def get_quant_method(self, prefix, module_kind): + return quant_method + + class FakeContext: + quant_config = FakeQuantConfig() + + monkeypatch.setattr(moe, 'get_build_model_context', lambda: FakeContext()) + + with pytest.raises(NotImplementedError, match='fp32_acc or output_scale'): + moe.build_fused_moe(16, + 32, + 4, + 2, + quant_config={}, + fp32_acc=True) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason='requires CUDA') +def test_moe_reduce_accumulates_in_fp32_before_scaling_and_casting(): + from lmdeploy.pytorch.kernels.cuda.moe.fused_moe import moe_reduce + + hidden = torch.tensor([[[1.25, -2.5], [3.0, 4.0]]], device='cuda', dtype=torch.bfloat16) + weights = torch.tensor([[0.2, 0.7]], device='cuda', dtype=torch.float32) + actual = moe_reduce(hidden, weights, fp32_acc=True, output_scale=2.5) + expected = ((hidden.float() * weights[..., None]).sum(dim=1) * 2.5).to(hidden.dtype) + + torch.testing.assert_close(actual, expected, rtol=0, atol=0) From 96a10e624624809b77059c4dde2a79d6db8dc7ef Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Wed, 23 Sep 2026 09:53:28 +0000 Subject: [PATCH 08/39] fix(pytorch): address GLM backend and runtime policy review Reject unsupported DLINFER MoE reduction options instead of ignoring them. Leave NCCL NVLS policy to the runtime and remove model-specific environment plumbing. Document GLM FP32 operator usage and cover dtype, reduction, and environment contracts. Validation: 110 focused tests passed on CPU/H200; Ruff passed for lmdeploy and tests/pytorch. DLINFER rejection regressions failed before the fix and pass afterward. --- lmdeploy/pytorch/backends/dlinfer/moe.py | 3 + lmdeploy/pytorch/config.py | 4 -- lmdeploy/pytorch/configurations/glm5_next.py | 4 -- .../pytorch/engine/executor/base_worker.py | 4 -- .../pytorch/engine/executor/ray_executor.py | 14 +---- lmdeploy/pytorch/nn/norm.py | 1 + lmdeploy/pytorch/nn/rotary_embedding.py | 6 +- tests/pytorch/config/test_glm5_runtime_env.py | 43 ++++++++++++++ tests/pytorch/nn/test_fp32_norm.py | 29 ++++++++++ tests/pytorch/nn/test_moe_options.py | 58 ++++++++++++++++++- tests/pytorch/nn/test_rotary_embedding.py | 31 ++++++++++ 11 files changed, 171 insertions(+), 26 deletions(-) create mode 100644 tests/pytorch/config/test_glm5_runtime_env.py create mode 100644 tests/pytorch/nn/test_fp32_norm.py diff --git a/lmdeploy/pytorch/backends/dlinfer/moe.py b/lmdeploy/pytorch/backends/dlinfer/moe.py index 166a36d04c..0cbd34a546 100644 --- a/lmdeploy/pytorch/backends/dlinfer/moe.py +++ b/lmdeploy/pytorch/backends/dlinfer/moe.py @@ -117,6 +117,9 @@ def forward(self, def _build_fused_moe(spec: FusedMoEBuildSpec) -> FusedMoEImpl: """Build a DLINFER fused MoE implementation.""" + if spec.fp32_acc or spec.output_scale != 1.0: + raise NotImplementedError( + 'DLINFER fused MoE does not support fp32_acc or output_scale.') return DlinferFusedMoEImpl( top_k=spec.top_k, num_experts=spec.num_experts, diff --git a/lmdeploy/pytorch/config.py b/lmdeploy/pytorch/config.py index a0dd0dc7f8..5a2e1c80ae 100644 --- a/lmdeploy/pytorch/config.py +++ b/lmdeploy/pytorch/config.py @@ -485,10 +485,6 @@ class ModelConfig: # Number of contiguous TP ranks that own the same logical KV-head shard. num_replicate_key_value_heads: int = 1 - # Model-specific defaults that must be present before the distributed - # process group is initialized. Explicit process environment values win. - process_group_env_defaults: dict[str, str] = field(default_factory=dict) - @property def use_mla_fp8_cache(self): """Whether MLA uses the DeepSeek-V3.2 FP8 cache layout.""" diff --git a/lmdeploy/pytorch/configurations/glm5_next.py b/lmdeploy/pytorch/configurations/glm5_next.py index fda90cf426..0982a53f06 100644 --- a/lmdeploy/pytorch/configurations/glm5_next.py +++ b/lmdeploy/pytorch/configurations/glm5_next.py @@ -222,10 +222,6 @@ def build(cls, hf_config, model_path: str | None = None, **kwargs): **dict(kwargs, is_draft_model=False)) tp = kwargs.get('tp', 1) - device_type = kwargs.get('device_type', 'auto') - if device_type == 'cuda' and tp > 1: - config.process_group_env_defaults.setdefault( - 'NCCL_NVLS_ENABLE', '0') num_linear_layers = len(linear_layer_ids) num_full_layers = len(full_attention_layer_ids) num_heads = linear_config['num_heads'] diff --git a/lmdeploy/pytorch/engine/executor/base_worker.py b/lmdeploy/pytorch/engine/executor/base_worker.py index 558afb84d8..0fc79182e7 100644 --- a/lmdeploy/pytorch/engine/executor/base_worker.py +++ b/lmdeploy/pytorch/engine/executor/base_worker.py @@ -1,7 +1,6 @@ # Copyright (c) OpenMMLab. All rights reserved. import asyncio import gc -import os from typing import Any from lmdeploy.pytorch.backends.selector import get_backend @@ -59,9 +58,6 @@ def __init__( def init_process_group(self, rank: int, master_addr: str = None, master_port: str = None): """Initialize process group.""" - for key, value in self.model_config.process_group_env_defaults.items(): - os.environ.setdefault(key, value) - self.rank = rank if self.world_size > 1: if master_addr is not None and master_port is not None: diff --git a/lmdeploy/pytorch/engine/executor/ray_executor.py b/lmdeploy/pytorch/engine/executor/ray_executor.py index c44ec3cb1c..3c77d91590 100644 --- a/lmdeploy/pytorch/engine/executor/ray_executor.py +++ b/lmdeploy/pytorch/engine/executor/ray_executor.py @@ -73,16 +73,11 @@ def _update_env_cuda_alloc_conf(env_vars: dict): env_vars['PYTORCH_CUDA_ALLOC_CONF'] = cuda_alloc_conf -def _update_runtime_envs( - runtime_env: dict, - process_group_env_defaults: dict[str, str] | None = None, -): +def _update_runtime_envs(runtime_env: dict): """Update runtime envs.""" new_envs = _envs.get_all_envs() env_vars: dict = runtime_env.get('env_vars', {}) env_vars.update(new_envs) - for key, default in (process_group_env_defaults or {}).items(): - env_vars.setdefault(key, os.environ.get(key, default)) _update_env_cuda_alloc_conf(env_vars) runtime_env['env_vars'] = env_vars return runtime_env @@ -710,8 +705,6 @@ def get_priority(ip): def _init_workers_ray(self, placement_group: PlacementGroup, worker_kwargs: dict): """Init worker ray.""" device_str = get_device_str() - process_group_env_defaults = ( - worker_kwargs['model_config'].process_group_env_defaults) bundle_indices = [] if not _envs.ray_external_pg_bundles: for bundle_id, bundle in enumerate(placement_group.bundle_specs): @@ -748,7 +741,7 @@ def _init_workers_ray(self, placement_group: PlacementGroup, worker_kwargs: dict if device_str == 'GPU': runtime_env = dict() - runtime_env = _update_runtime_envs(runtime_env, process_group_env_defaults) + runtime_env = _update_runtime_envs(runtime_env) if self._needs_symm_mem_device_setup: # Symmetric-memory IPC needs peer TP GPUs to stay visible. # Keep the inherited visibility and bind each actor below. @@ -763,8 +756,7 @@ def _init_workers_ray(self, placement_group: PlacementGroup, worker_kwargs: dict )(RayWorkerWrapper).remote(**worker_kwargs) else: runtime_env = dict() - runtime_env = _update_runtime_envs( - runtime_env, process_group_env_defaults) + runtime_env = _update_runtime_envs(runtime_env) worker = ray.remote( num_cpus=0, num_gpus=0, diff --git a/lmdeploy/pytorch/nn/norm.py b/lmdeploy/pytorch/nn/norm.py index b39ed8d75a..62b9bf5552 100644 --- a/lmdeploy/pytorch/nn/norm.py +++ b/lmdeploy/pytorch/nn/norm.py @@ -34,6 +34,7 @@ class FP32LayerNorm(nn.Module): Some model components keep LayerNorm weights in FP32 even when the model activation dtype is BF16. Keep that numerical contract in one reusable module and cast only the returned activation back to its input dtype. + Used by GLM-5.3 vision patch merging and KPool key normalization. """ def __init__(self, diff --git a/lmdeploy/pytorch/nn/rotary_embedding.py b/lmdeploy/pytorch/nn/rotary_embedding.py index 68452d47f0..b157378b62 100644 --- a/lmdeploy/pytorch/nn/rotary_embedding.py +++ b/lmdeploy/pytorch/nn/rotary_embedding.py @@ -256,7 +256,11 @@ def apply_rotary_pos_emb_fp32(query: Tensor, cos: Tensor, sin: Tensor, unsqueeze_dim: int = 1) -> tuple[Tensor, Tensor]: - """Apply NeoX-style RoPE with FP32 arithmetic and dtype-preserving output.""" + """Apply NeoX-style RoPE with FP32 arithmetic and dtype-preserving output. + + Used by GLM-5.3 vision attention to match its FP32 rotary arithmetic; + the text attention uses NoPE and does not call this helper. + """ query_dtype = query.dtype key_dtype = key.dtype query = query.float() diff --git a/tests/pytorch/config/test_glm5_runtime_env.py b/tests/pytorch/config/test_glm5_runtime_env.py new file mode 100644 index 0000000000..61bf6803ba --- /dev/null +++ b/tests/pytorch/config/test_glm5_runtime_env.py @@ -0,0 +1,43 @@ +# Copyright (c) OpenMMLab. All rights reserved. +import os +from unittest.mock import Mock + +import pytest + +from lmdeploy.hf_configs.configuration_glm5_next import Glm5NextConfig +from lmdeploy.pytorch.config import DistConfig +from lmdeploy.pytorch.configurations.glm5_next import Glm5NextModelConfigBuilder +from lmdeploy.pytorch.engine.executor import base_worker + + +@pytest.mark.parametrize('nvls', [None, '0', '1']) +def test_glm5_leaves_nvls_policy_to_runtime(monkeypatch, nvls): + if nvls is None: + monkeypatch.delenv('NCCL_NVLS_ENABLE', raising=False) + else: + monkeypatch.setenv('NCCL_NVLS_ENABLE', nvls) + hf_config = Glm5NextConfig(text_config={ + 'num_hidden_layers': 4, + 'layer_types': ['linear_attention'] * 3 + ['deepseek_sparse_attention'], + 'linear_num_heads': 64, + 'linear_head_dim': 128, + 'linear_conv_kernel_dim': 4, + 'index_kpool': 4, + }) + monkeypatch.setattr('lmdeploy.pytorch.configurations.deepseek_v2.flash_mla_available', lambda: True) + config = Glm5NextModelConfigBuilder.build(hf_config, tp=8, device_type='cuda') + worker = base_worker.WorkerWrapperBase.__new__(base_worker.WorkerWrapperBase) + worker.model_config = config + worker.dist_config = DistConfig(tp=8) + worker.world_size = 8 + worker.device_type = 'cuda' + seen = [] + monkeypatch.setattr(base_worker, 'init_process_group', + lambda rank, size: seen.append(os.environ.get('NCCL_NVLS_ENABLE'))) + monkeypatch.setattr(base_worker, 'get_backend', Mock()) + monkeypatch.setattr(base_worker.DistContext, 'build', Mock()) + + worker.init_process_group(rank=0) + + assert seen == [nvls] + assert os.environ.get('NCCL_NVLS_ENABLE') == nvls diff --git a/tests/pytorch/nn/test_fp32_norm.py b/tests/pytorch/nn/test_fp32_norm.py new file mode 100644 index 0000000000..52c745d763 --- /dev/null +++ b/tests/pytorch/nn/test_fp32_norm.py @@ -0,0 +1,29 @@ +# Copyright (c) OpenMMLab. All rights reserved. +import pytest +import torch +import torch.nn.functional as F + +from lmdeploy.pytorch.nn.norm import FP32LayerNorm + + +@pytest.mark.parametrize('device', [ + 'cpu', pytest.param('cuda', marks=pytest.mark.skipif( + not torch.cuda.is_available(), reason='requires CUDA')), +]) +@pytest.mark.parametrize('dtype', [torch.float16, torch.bfloat16, torch.float32]) +@pytest.mark.parametrize('bias', [False, True]) +def test_fp32_layer_norm_preserves_parameters_and_output_dtype(dtype, bias, device): + norm = FP32LayerNorm(16, bias=bias, device=device) + generator = torch.Generator().manual_seed(123) + inputs = torch.randn(2, 3, 16, generator=generator).to(device=device, dtype=dtype) + norm.weight.data.copy_(torch.randn(16, generator=generator)) + if bias: + norm.bias.data.copy_(torch.randn(16, generator=generator)) + + expected = F.layer_norm(inputs.float(), (16,), norm.weight, norm.bias, norm.eps).to(dtype) + actual = norm(inputs) + + assert norm.weight.dtype == torch.float32 + assert norm.bias is None or norm.bias.dtype == torch.float32 + assert actual.dtype == dtype + torch.testing.assert_close(actual, expected, rtol=0, atol=0) diff --git a/tests/pytorch/nn/test_moe_options.py b/tests/pytorch/nn/test_moe_options.py index 7f88b6887a..49f5d07347 100644 --- a/tests/pytorch/nn/test_moe_options.py +++ b/tests/pytorch/nn/test_moe_options.py @@ -1,4 +1,9 @@ # Copyright (c) OpenMMLab. All rights reserved. +import importlib.util +import sys +from types import ModuleType +from unittest.mock import Mock + import pytest import torch @@ -28,8 +33,12 @@ def __init__(self, **kwargs): assert captured['output_scale'] == 2.5 +@pytest.mark.parametrize('options', [ + {'fp32_acc': True}, {'output_scale': 2.5}, + {'fp32_acc': True, 'output_scale': 2.5}, +]) @pytest.mark.parametrize('quant_method', ['smooth_quant', 'fp8', 'compressed-tensors']) -def test_build_fused_moe_rejects_unsupported_reduction_options(monkeypatch, quant_method): +def test_build_fused_moe_rejects_unsupported_reduction_options(monkeypatch, quant_method, options): from lmdeploy.pytorch.nn import moe class FakeQuantConfig: @@ -53,7 +62,7 @@ class FakeContext: 4, 2, quant_config={}, - fp32_acc=True) + **options) @pytest.mark.skipif(not torch.cuda.is_available(), reason='requires CUDA') @@ -66,3 +75,48 @@ def test_moe_reduce_accumulates_in_fp32_before_scaling_and_casting(): expected = ((hidden.float() * weights[..., None]).sum(dim=1) * 2.5).to(hidden.dtype) torch.testing.assert_close(actual, expected, rtol=0, atol=0) + + +@pytest.fixture +def dlinfer_moe(monkeypatch): + # Exercise the real builder without requiring a DLINFER vendor runtime. + import lmdeploy.pytorch.backends.dlinfer as backend + + kernels = ModuleType('lmdeploy.pytorch.kernels.dlinfer') + for name in ('DlinferMoECommType', 'DlinferMoeMetadata', 'fused_moe', + 'fused_moe_w8a8', 'moe_gating_topk_softmax'): + setattr(kernels, name, Mock()) + spec = importlib.util.spec_from_file_location( + f'{backend.__name__}._test_moe', f'{backend.__path__[0]}/moe.py') + module = importlib.util.module_from_spec(spec) + with monkeypatch.context() as patch: + patch.setitem(sys.modules, kernels.__name__, kernels) + spec.loader.exec_module(module) + return module + + +def _dlinfer_build_spec(**options): + from lmdeploy.pytorch.backends.moe import FusedMoEBuildSpec + + return FusedMoEBuildSpec(top_k=2, num_experts=4, renormalize=True, + hidden_dim=16, ep_size=1, ep_group=None, + layer_idx=0, output_dtype=torch.bfloat16, + num_max_dispatch_tokens_per_rank=32, **options) + + +def test_dlinfer_moe_accepts_default_reduction_options(dlinfer_moe): + impl = dlinfer_moe._build_fused_moe(_dlinfer_build_spec()) + assert isinstance(impl, dlinfer_moe.DlinferFusedMoEImpl) + assert (impl.top_k, impl.num_experts, impl.renormalize, impl.ep_size) == (2, 4, True, 1) + + +@pytest.mark.parametrize('options', [ + {'fp32_acc': True}, {'output_scale': 2.5}, + {'fp32_acc': True, 'output_scale': 2.5}, +]) +def test_dlinfer_moe_rejects_unsupported_reduction_options(dlinfer_moe, monkeypatch, options): + constructor = Mock() + monkeypatch.setattr(dlinfer_moe, 'DlinferFusedMoEImpl', constructor) + with pytest.raises(NotImplementedError, match='fp32_acc or output_scale'): + dlinfer_moe._build_fused_moe(_dlinfer_build_spec(**options)) + constructor.assert_not_called() diff --git a/tests/pytorch/nn/test_rotary_embedding.py b/tests/pytorch/nn/test_rotary_embedding.py index 00964fac86..cea8bb3892 100644 --- a/tests/pytorch/nn/test_rotary_embedding.py +++ b/tests/pytorch/nn/test_rotary_embedding.py @@ -1,3 +1,4 @@ +import pytest import torch from transformers import PretrainedConfig @@ -226,3 +227,33 @@ def test_default_apply_rotary_complex_accepts_half_width_tables_with_empty_key() torch.testing.assert_close(q_embed, _complex_rope_reference(q_states, cos, sin)) assert k_embed.shape == k_states.shape + + +@pytest.mark.parametrize('device', [ + 'cpu', pytest.param('cuda', marks=pytest.mark.skipif( + not torch.cuda.is_available(), reason='requires CUDA')), +]) +@pytest.mark.parametrize('dtype', [torch.float16, torch.bfloat16, torch.float32]) +@pytest.mark.parametrize('unsqueeze_dim', [0, 1]) +def test_fp32_rotary_matches_reference_and_preserves_input(dtype, unsqueeze_dim, device): + from lmdeploy.pytorch.nn.rotary_embedding import apply_rotary_pos_emb_fp32 + + generator = torch.Generator().manual_seed(123) + shape = (3, 5, 16) if unsqueeze_dim == 0 else (5, 3, 16) + query = torch.randn(shape, generator=generator).to(device=device, dtype=dtype) + key = torch.randn(shape, generator=generator).to(device=device, dtype=dtype) + # Exactly representable coefficients isolate intermediate rounding in BF16/FP16. + cos = torch.full((5, 16), 0.625, dtype=dtype, device=device) + sin = torch.full((5, 16), 0.375, dtype=dtype, device=device) + original = (query.clone(), key.clone()) + outputs = apply_rotary_pos_emb_fp32(query, key, cos, sin, unsqueeze_dim) + + for value, saved, actual in zip((query, key), original, outputs): + left, right = value.float().chunk(2, dim=-1) + expected = torch.cat((left * 0.625 - right * 0.375, + right * 0.625 + left * 0.375), dim=-1).to(dtype) + assert actual.dtype == dtype + torch.testing.assert_close(actual, expected, + rtol=1e-6 if dtype == torch.float32 else 0, + atol=1e-7 if dtype == torch.float32 else 0) + torch.testing.assert_close(value, saved, rtol=0, atol=0) From a17e90b4f529bb3e56e8c9276cec169791bf17ff Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Wed, 23 Sep 2026 10:12:02 +0000 Subject: [PATCH 09/39] style(pytorch): satisfy GLM pre-commit hooks Apply the configured string and docstring formatters to GLM PR files. Verified executable AST is unchanged; both previously failing hooks and Ruff pass locally. --- lmdeploy/pytorch/backends/cuda/kda.py | 17 +++---- lmdeploy/pytorch/backends/cuda/kpool.py | 6 ++- lmdeploy/pytorch/engine/logits_process.py | 4 +- .../kernels/cuda/sparse_mla_tilelang.py | 25 +++++----- lmdeploy/pytorch/models/glm5_next.py | 10 ++-- lmdeploy/pytorch/nn/attention.py | 3 +- lmdeploy/pytorch/nn/kpool.py | 50 +++++++++---------- lmdeploy/pytorch/nn/norm.py | 7 ++- lmdeploy/pytorch/nn/rotary_embedding.py | 4 +- 9 files changed, 61 insertions(+), 65 deletions(-) diff --git a/lmdeploy/pytorch/backends/cuda/kda.py b/lmdeploy/pytorch/backends/cuda/kda.py index ff3511cec1..f7603c11b4 100644 --- a/lmdeploy/pytorch/backends/cuda/kda.py +++ b/lmdeploy/pytorch/backends/cuda/kda.py @@ -1,9 +1,8 @@ # Copyright (c) OpenMMLab. All rights reserved. """CUDA KDA backend composed from FLA and shared LMDeploy operators. -KDA is distinct from LMDeploy's gated-delta rule, but its CUDA implementation -does not need copied GLM kernels. This adapter owns only LMDeploy cache/state -semantics, reuses FLA convolution/prefill, and shares the TileLang recurrent +KDA is distinct from LMDeploy's gated-delta rule, but its CUDA implementation does not need copied GLM kernels. This +adapter owns only LMDeploy cache/state semantics, reuses FLA convolution/prefill, and shares the TileLang recurrent state-ring kernel with gated-delta rule for both AR and MTP decode. """ @@ -68,9 +67,8 @@ def _forward_spec(self, mixed_qkv, raw_gate, raw_beta, conv_state, recurrent_state, metadata, **kwargs): """Checkpoint every verified token at its accepted-history ring slot. - State is addressed by accepted history length, not the last proposed - length. The ring therefore also handles zero/partial acceptance and - request reordering without a scheduler-side rollback hook. + State is addressed by accepted history length, not the last proposed length. The ring therefore also handles + zero/partial acceptance and request reordering without a scheduler-side rollback hook. """ if metadata.is_decoding: return self._forward_spec_decode(mixed_qkv, raw_gate, raw_beta, @@ -117,10 +115,11 @@ def _decode_recurrent(self, q, k, v, g, beta, A_log, dt_bias, initial_state, def _forward_spec_decode(self, mixed_qkv, raw_gate, raw_beta, conv_state, recurrent_state, metadata, **kwargs): - """Reuse causal-convolution token rings and one verification recurrence. + """Reuse causal-convolution token rings and one verification + recurrence. - The recurrence is parallel across state tiles, not across causally - dependent timesteps. Each timestep is saved for partial acceptance. + The recurrence is parallel across state tiles, not across causally dependent timesteps. Each timestep is saved + for partial acceptance. """ ids = metadata.state_ids.long() batch = ids.numel() diff --git a/lmdeploy/pytorch/backends/cuda/kpool.py b/lmdeploy/pytorch/backends/cuda/kpool.py index 2fd78fab41..400526f32f 100644 --- a/lmdeploy/pytorch/backends/cuda/kpool.py +++ b/lmdeploy/pytorch/backends/cuda/kpool.py @@ -43,7 +43,8 @@ def kpool_compress_quantize_cuda( mode: str, round_scale: bool, ) -> tuple[Tensor, Tensor]: - """Compress and quantize closed pools with LMDeploy's reusable semantics.""" + """Compress and quantize closed pools with LMDeploy's reusable + semantics.""" pooled = kpool_compress(slot_k, slot_score, ape, mode=mode) return kpool_quantize_fp8( pooled, @@ -60,7 +61,8 @@ def kpool_select_groups_cuda( row_starts: Tensor | None = None, max_group_length: int | None = None, ) -> Tensor: - """Select pooled groups with LMDeploy's shared sparse-index Top-K kernel.""" + """Select pooled groups with LMDeploy's shared sparse-index Top-K + kernel.""" if not is_sparse_index_topk_supported(group_topk): raise ValueError( 'The GLM-5.3 KPool selector only supports group_topk=512 or 2048, ' diff --git a/lmdeploy/pytorch/engine/logits_process.py b/lmdeploy/pytorch/engine/logits_process.py index ed74a94d6b..30853ff882 100644 --- a/lmdeploy/pytorch/engine/logits_process.py +++ b/lmdeploy/pytorch/engine/logits_process.py @@ -522,8 +522,8 @@ def _filter_sorted_logits(self, logits: torch.Tensor): def filter_logits(self, logits: torch.Tensor): """Apply sampling filters in vocabulary order without modifying logits. - Speculative verification needs the same target distribution as AR, - but its backend consumes logits rather than sampled token IDs. + Speculative verification needs the same target distribution as AR, but its backend consumes logits rather than + sampled token IDs. """ inputs = self.sampling_inputs if inputs.max_top_k <= 0 and inputs.top_k is None and inputs.top_p is None and inputs.min_p is None: diff --git a/lmdeploy/pytorch/kernels/cuda/sparse_mla_tilelang.py b/lmdeploy/pytorch/kernels/cuda/sparse_mla_tilelang.py index 3328b70b49..247d8876cd 100644 --- a/lmdeploy/pytorch/kernels/cuda/sparse_mla_tilelang.py +++ b/lmdeploy/pytorch/kernels/cuda/sparse_mla_tilelang.py @@ -2,9 +2,8 @@ # SPDX-License-Identifier: Apache-2.0 """TileLang sparse MLA decode attention for CUDA BF16 tensors. -The kernel schedule in this file is adapted from SGLang's -sparse_attention_fwd_kernel_v1 at commit 9e692c9216c3, distributed under the -Apache License 2.0: +The kernel schedule in this file is adapted from SGLang's sparse_attention_fwd_kernel_v1 at commit 9e692c9216c3, +distributed under the Apache License 2.0: https://github.com/sgl-project/sglang/blob/9e692c9216c3/python/sglang/kernels/ops/attention/dsa/tilelang_kernel.py @@ -45,27 +44,27 @@ def _sparse_mla_bf16_fwd_kernel( assert ( dim == tilelang.math.next_power_of_2(dim) or dim % 64 == 0 ), f"dim={dim} must be a power of 2 or a multiple of 64" - assert is_causal, "non-causal is not supported" + assert is_causal, 'non-causal is not supported' assert ( topk % block_I == 0 - ), "otherwise will load some index=0 thus causing wrong kv to be loaded" + ), 'otherwise will load some index=0 thus causing wrong kv to be loaded' if sm_scale is None: sm_scale = (1.0 / dim) ** 0.5 * 1.44269504 # log2(e) else: sm_scale = sm_scale * 1.44269504 # log2(e) - batch = T.symbolic("batch") - seq_len = T.symbolic("seq_len") - seq_len_kv = T.symbolic("seq_len_kv") + batch = T.symbolic('batch') + seq_len = T.symbolic('seq_len') + seq_len_kv = T.symbolic('seq_len_kv') head_kv = num_heads // kv_group q_shape = [batch, seq_len, num_heads, dim] kv_shape = [batch, seq_len_kv, kv_group, storage_dim] o_shape = [batch, seq_len, num_heads, dim] indices_shape = [batch, seq_len, kv_group, topk] - indices_dtype = "int32" - dtype = "bfloat16" - accum_dtype = "float" + indices_dtype = 'int32' + dtype = 'bfloat16' + accum_dtype = 'float' H = head_kv padded_H = max(tilelang.math.next_power_of_2(head_kv), 16) @@ -76,7 +75,7 @@ def _sparse_mla_bf16_fwd_kernel( D = dim if head_kv > 64: - assert head_kv % 64 == 0, "head_kv should be a multiple of 64" + assert head_kv % 64 == 0, 'head_kv should be a multiple of 64' REPLICATE_H = head_kv // 64 else: REPLICATE_H = 1 @@ -98,7 +97,7 @@ def main( Q_shared = T.alloc_shared([H_per_block, D], dtype) KV_shared = T.alloc_shared([BI, D], dtype) O_shared = T.alloc_shared([H_per_block, D], dtype) - mask = T.alloc_fragment([BI], "bool") + mask = T.alloc_fragment([BI], 'bool') acc_o = T.alloc_fragment([H_per_block, D], accum_dtype) acc_s = T.alloc_fragment([H_per_block, BI], accum_dtype) diff --git a/lmdeploy/pytorch/models/glm5_next.py b/lmdeploy/pytorch/models/glm5_next.py index 5a184a28dc..5fbd933425 100644 --- a/lmdeploy/pytorch/models/glm5_next.py +++ b/lmdeploy/pytorch/models/glm5_next.py @@ -1923,10 +1923,9 @@ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]], *, class Glm5NextMTPAttention(Glm5NextSparseAttention): """The predictor's KPool tail is reconstructed from pageable token data. - Draft forwards can revisit accepted positions after multiple proposals. - Keeping raw index keys/scores in its cache avoids private mutable request - state and reuses the normal cache allocation, sizing and sleep lifecycle. - Only the single MTP layer requests this additional cache. + Draft forwards can revisit accepted positions after multiple proposals. Keeping raw index keys/scores in its cache + avoids private mutable request state and reuses the normal cache allocation, sizing and sleep lifecycle. Only the + single MTP layer requests this additional cache. """ _TOKEN_CACHE = 'glm5_mtp_kpool_tokens' @@ -2003,7 +2002,8 @@ def forward(self, hidden_states, rotary_pos_emb, past_key_value, class Glm5NextMTPModel(GlmMoeDsaMTPModel): - """Reuse the shared GLM/DeepSeek predictor, proposer and CUDA Graph flow.""" + """Reuse the shared GLM/DeepSeek predictor, proposer and CUDA Graph + flow.""" uses_shared_input_embeddings = True diff --git a/lmdeploy/pytorch/nn/attention.py b/lmdeploy/pytorch/nn/attention.py index 3edab3e6f2..0d89e25a7e 100644 --- a/lmdeploy/pytorch/nn/attention.py +++ b/lmdeploy/pytorch/nn/attention.py @@ -102,7 +102,8 @@ def fill_and_flatten_latent_kv_cache( k_scales_zeros: torch.Tensor = None, v_scales_zeros: torch.Tensor = None, ) -> torch.Tensor: - """Append latent KV and return the complete request-major prefill KV.""" + """Append latent KV and return the complete request-major prefill + KV.""" self._lazy_init(key.device) quant_policy = attn_metadata.quant_policy diff --git a/lmdeploy/pytorch/nn/kpool.py b/lmdeploy/pytorch/nn/kpool.py index 83eb6fcd28..bd3ee3fc68 100644 --- a/lmdeploy/pytorch/nn/kpool.py +++ b/lmdeploy/pytorch/nn/kpool.py @@ -1,14 +1,12 @@ # Copyright (c) OpenMMLab. All rights reserved. """Reusable Torch semantics for the DSA KPool indexer. -KPool has two different kinds of runtime data. Closed pools are pageable and -belong in the named DSA index cache. The unfinished per-request tail is -sequence state and must be supplied by the caller; this module deliberately -does not hide it in mutable module tensors. - -The functions here are a device-agnostic correctness path. CUDA backends can -replace compression, FP8 scoring, and top-k with fused kernels while retaining -these input/output contracts. +KPool has two different kinds of runtime data. Closed pools are pageable and belong in the named DSA index cache. The +unfinished per-request tail is sequence state and must be supplied by the caller; this module deliberately does not hide +it in mutable module tensors. + +The functions here are a device-agnostic correctness path. CUDA backends can replace compression, FP8 scoring, and +top-k with fused kernels while retaining these input/output contracts. """ from __future__ import annotations @@ -90,9 +88,8 @@ class KPoolDecodeUpdate: class KPoolIndexer(nn.Module): """Replicated KPool parameter layer shared by all attention-TP ranks. - The seven parameter names and dtypes match the GLM-5.3/SGLang checkpoint - contract. Query/key rotary handling and cache ownership stay with the - model/backend because they depend on model geometry and request metadata. + The seven parameter names and dtypes match the GLM-5.3/SGLang checkpoint contract. Query/key rotary handling and + cache ownership stay with the model/backend because they depend on model geometry and request metadata. """ def __init__( @@ -253,7 +250,8 @@ def kpool_normalized_hadamard(values: Tensor) -> Tensor: def kpool_rotate_query(query: Tensor) -> Tensor: - """Match the BF16-preserving query rotation used before FP8 quantization.""" + """Match the BF16-preserving query rotation used before FP8 + quantization.""" return kpool_normalized_hadamard(query).to(query.dtype) @@ -267,9 +265,8 @@ def kpool_partition_update( ) -> KPoolUpdate: """Assemble arbitrary-length input with a prior tail into closed pools. - This is the state-free equivalent of SGLang's extend/decode tail ring. The - caller owns persistence of the returned tail and supplies it on the next - invocation. + This is the state-free equivalent of SGLang's extend/decode tail ring. The caller owns persistence of the returned + tail and supplies it on the next invocation. """ _validate_pool_geometry(pool_size) if history_length < 0: @@ -328,10 +325,9 @@ def kpool_decode_update( ) -> KPoolDecodeUpdate: """Build a fixed-shape, graph-safe update for one-token decode. - State slot zero is LMDeploy's reserved dummy slot. Invalid CUDA Graph - padding rows therefore read and write slot zero without affecting a live - request. Every returned tensor has a shape determined only by the graph - capture bucket. + State slot zero is LMDeploy's reserved dummy slot. Invalid CUDA Graph padding rows therefore read and write slot + zero without affecting a live request. Every returned tensor has a shape determined only by the graph capture + bucket. """ _validate_pool_geometry(pool_size) if keys.ndim != 2 or scores.shape != keys.shape: @@ -420,9 +416,8 @@ def kpool_compress_online(closed_keys: Tensor, closed_scores: Tensor, ape: Tensor) -> Tensor: """Use the online recurrence from SGLang's extend assembly kernel. - The explicit slot loop matches SGLang's compression kernel reduction - order. A vectorized max/exp/sum is mathematically equivalent but can move - values across the FP8 quantization boundary after different FP32 rounds. + The explicit slot loop matches SGLang's compression kernel reduction order. A vectorized max/exp/sum is + mathematically equivalent but can move values across the FP8 quantization boundary after different FP32 rounds. """ _validate_compress_inputs(closed_keys, closed_scores, ape) @@ -550,7 +545,8 @@ def kpool_score( def kpool_pooled_block_offsets(token_block_offsets: Tensor, pool_size: int) -> Tensor: - """Build the pooled page table by selecting every ``pool_size`` token page.""" + """Build the pooled page table by selecting every ``pool_size`` token + page.""" _validate_pool_geometry(pool_size) if token_block_offsets.ndim < 1: raise ValueError('token_block_offsets must have at least one dimension.') @@ -655,9 +651,8 @@ def kpool_write_packed_cache_batched( ) -> None: """Write at most one closed pool per fixed-shape decode row. - Invalid rows target reserved cache block zero and keep its previous value. - This avoids dynamic boolean indexing and keeps CUDA Graph addresses and - launch geometry stable. + Invalid rows target reserved cache block zero and keep its previous value. This avoids dynamic boolean indexing and + keeps CUDA Graph addresses and launch geometry stable. """ if pooled_key_fp8.ndim != 2: raise ValueError( @@ -722,7 +717,8 @@ def kpool_read_packed_cache( def kpool_selected_token_counts(seq_lens: Tensor, topk: int, pool_size: int) -> Tensor: - """Return selected history plus always-selected ragged tail token counts.""" + """Return selected history plus always-selected ragged tail token + counts.""" _validate_pool_geometry(pool_size, topk) full_pool_tokens = torch.div(seq_lens, pool_size, rounding_mode='floor') * pool_size return full_pool_tokens.clamp(max=topk) + seq_lens - full_pool_tokens diff --git a/lmdeploy/pytorch/nn/norm.py b/lmdeploy/pytorch/nn/norm.py index 62b9bf5552..2ba7554969 100644 --- a/lmdeploy/pytorch/nn/norm.py +++ b/lmdeploy/pytorch/nn/norm.py @@ -31,10 +31,9 @@ def rms_scale(a: torch.Tensor, b: torch.Tensor, dim: int = -1, eps: float = 1e-6 class FP32LayerNorm(nn.Module): """LayerNorm with FP32 parameters and accumulation. - Some model components keep LayerNorm weights in FP32 even when the model - activation dtype is BF16. Keep that numerical contract in one reusable - module and cast only the returned activation back to its input dtype. - Used by GLM-5.3 vision patch merging and KPool key normalization. + Some model components keep LayerNorm weights in FP32 even when the model activation dtype is BF16. Keep that + numerical contract in one reusable module and cast only the returned activation back to its input dtype. Used by + GLM-5.3 vision patch merging and KPool key normalization. """ def __init__(self, diff --git a/lmdeploy/pytorch/nn/rotary_embedding.py b/lmdeploy/pytorch/nn/rotary_embedding.py index b157378b62..e94ce43de3 100644 --- a/lmdeploy/pytorch/nn/rotary_embedding.py +++ b/lmdeploy/pytorch/nn/rotary_embedding.py @@ -258,8 +258,8 @@ def apply_rotary_pos_emb_fp32(query: Tensor, unsqueeze_dim: int = 1) -> tuple[Tensor, Tensor]: """Apply NeoX-style RoPE with FP32 arithmetic and dtype-preserving output. - Used by GLM-5.3 vision attention to match its FP32 rotary arithmetic; - the text attention uses NoPE and does not call this helper. + Used by GLM-5.3 vision attention to match its FP32 rotary arithmetic; the text attention uses NoPE and does not call + this helper. """ query_dtype = query.dtype key_dtype = key.dtype From 0f330609c2316ce59ab37b0725e6ef28ef44fcbb Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Wed, 23 Sep 2026 11:27:13 +0000 Subject: [PATCH 10/39] refactor(pytorch): reuse native LayerNorm for GLM --- lmdeploy/pytorch/models/glm5_next.py | 21 ++------ lmdeploy/pytorch/nn/__init__.py | 2 +- lmdeploy/pytorch/nn/kpool.py | 9 ++-- lmdeploy/pytorch/nn/norm.py | 39 -------------- tests/pytorch/nn/test_fp32_norm.py | 29 ----------- tests/pytorch/nn/test_glm5_layer_norm.py | 65 ++++++++++++++++++++++++ 6 files changed, 75 insertions(+), 90 deletions(-) delete mode 100644 tests/pytorch/nn/test_fp32_norm.py create mode 100644 tests/pytorch/nn/test_glm5_layer_norm.py diff --git a/lmdeploy/pytorch/models/glm5_next.py b/lmdeploy/pytorch/models/glm5_next.py index 5fbd933425..c64a6d3798 100644 --- a/lmdeploy/pytorch/models/glm5_next.py +++ b/lmdeploy/pytorch/models/glm5_next.py @@ -34,7 +34,6 @@ from lmdeploy.pytorch.model_inputs import StepContext, StepContextManager, get_step_ctx_manager from lmdeploy.pytorch.nn import ( FlashAttention, - FP32LayerNorm, HcPrePost, Kda, KPoolIndexer, @@ -78,15 +77,6 @@ Glm5NextVisionRMSNorm = RMSNorm -class Glm5NextLayerNorm(FP32LayerNorm): - """GLM-5.3 FP32 LayerNorm using LMDeploy's reusable implementation.""" - - -# Backward-compatible name retained for the vision numerical contract tests -# and downstream imports. The provider is also shared by the KPool indexer. -Glm5NextVisionLayerNorm = Glm5NextLayerNorm - - def _build_glm53_latent_norm(hidden_size: int, eps: float, dtype: torch.dtype | None, device: torch.device | None) -> RMSNorm: @@ -240,9 +230,8 @@ def __init__(self, bias=False, dtype=dtype, device=device) - self.post_projection_norm = Glm5NextLayerNorm(dim, - eps=1e-6, - device=device) + self.post_projection_norm = nn.LayerNorm( + dim, eps=1e-6, dtype=torch.float32, device=device).requires_grad_(False) self.gate_up_proj = build_merged_colwise_linear( in_features=dim, all_out_features=[context_dim, context_dim], @@ -265,7 +254,8 @@ def __init__(self, def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: hidden_states = self.proj(hidden_states) - hidden_states = self.act1(self.post_projection_norm(hidden_states)) + hidden_states = self.post_projection_norm(hidden_states.float()).to(hidden_states.dtype) + hidden_states = self.act1(hidden_states) hidden_states = self.gate_up_proj(hidden_states) hidden_states = _glm_swiglu_impl(hidden_states, self.swiglu_limit, @@ -822,9 +812,6 @@ def _build_indexer(self, config: Any, layer_idx: int, dtype: torch.dtype, index_kpool=config.index_kpool, dtype=dtype, device=device, - key_norm=Glm5NextLayerNorm(config.index_head_dim, - eps=1e-6, - device=device), ) def _qkv_proj_unabsorbed(self, hidden_states: torch.Tensor, diff --git a/lmdeploy/pytorch/nn/__init__.py b/lmdeploy/pytorch/nn/__init__.py index f321e96cdd..f31de8f881 100644 --- a/lmdeploy/pytorch/nn/__init__.py +++ b/lmdeploy/pytorch/nn/__init__.py @@ -7,7 +7,7 @@ from .hc_prepost import HcPrePost # noqa: F401 from .kda import Kda # noqa: F401 from .kpool import KPoolIndexer # noqa: F401 -from .norm import FP32LayerNorm, LayerNorm, RMSNorm, rms_scale # noqa: F401 +from .norm import LayerNorm, RMSNorm, rms_scale # noqa: F401 from .rotary_embedding import ( ApplyRotaryEmb, # noqa: F401 RopeType, # noqa: F401 diff --git a/lmdeploy/pytorch/nn/kpool.py b/lmdeploy/pytorch/nn/kpool.py index bd3ee3fc68..80bd718766 100644 --- a/lmdeploy/pytorch/nn/kpool.py +++ b/lmdeploy/pytorch/nn/kpool.py @@ -22,7 +22,6 @@ from lmdeploy.pytorch.engine.cache_engine.schema import BlockCacheBinding, BlockCacheRequest, BlockCacheRequestContext from lmdeploy.pytorch.model_inputs import get_step_ctx_manager from lmdeploy.pytorch.nn.linear import build_colwise_linear -from lmdeploy.pytorch.nn.norm import FP32LayerNorm KPOOL_PAGE_SIZE = 64 KPOOL_FP8_MAX = 448.0 @@ -167,8 +166,8 @@ def add_prefix(name: str) -> str: # The device-agnostic KPool owner keeps a Torch reference default. # Models that require a platform-exact provider can inject the same # parameter-shaped component without changing checkpoint names. - self.k_norm = (key_norm if key_norm is not None else FP32LayerNorm( - index_head_dim, norm_eps, device=device)) + self.k_norm = (key_norm if key_norm is not None else nn.LayerNorm( + index_head_dim, eps=norm_eps, dtype=torch.float32, device=device).requires_grad_(False)) self.index_kpool_compress_ape = nn.Parameter( torch.zeros(index_kpool, index_head_dim, dtype=torch.float32, device=device), requires_grad=False, @@ -213,7 +212,9 @@ def project_query(self, q_lora: Tensor) -> Tensor: def project_key(self, hidden_states: Tensor) -> Tensor: """Project and normalize one shared index key per token.""" - return self.k_norm(self.wk(hidden_states)) + key = self.wk(hidden_states) + # Match vLLM's FP32 LayerNorm followed by a cast to the activation dtype. + return self.k_norm(key.float()).to(key.dtype) def project_compress_score(self, hidden_states: Tensor) -> Tensor: """Return per-slot, per-dimension compression gates.""" diff --git a/lmdeploy/pytorch/nn/norm.py b/lmdeploy/pytorch/nn/norm.py index 2ba7554969..362c2ef6e5 100644 --- a/lmdeploy/pytorch/nn/norm.py +++ b/lmdeploy/pytorch/nn/norm.py @@ -1,7 +1,6 @@ # Copyright (c) OpenMMLab. All rights reserved. import torch -import torch.nn.functional as F from torch import nn from lmdeploy.pytorch.distributed import get_dist_group, get_tp_world_rank @@ -28,44 +27,6 @@ def rms_scale(a: torch.Tensor, b: torch.Tensor, dim: int = -1, eps: float = 1e-6 return out.to(result_dtype) -class FP32LayerNorm(nn.Module): - """LayerNorm with FP32 parameters and accumulation. - - Some model components keep LayerNorm weights in FP32 even when the model activation dtype is BF16. Keep that - numerical contract in one reusable module and cast only the returned activation back to its input dtype. Used by - GLM-5.3 vision patch merging and KPool key normalization. - """ - - def __init__(self, - hidden_size: int, - eps: float = 1e-6, - bias: bool = True, - device: torch.device | str | None = None): - super().__init__() - self.hidden_size = hidden_size - self.eps = eps - self.weight = nn.Parameter(torch.ones(hidden_size, - dtype=torch.float32, - device=device), - requires_grad=False) - if bias: - self.bias = nn.Parameter(torch.zeros(hidden_size, - dtype=torch.float32, - device=device), - requires_grad=False) - else: - self.register_parameter('bias', None) - - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - """Normalize in FP32 and preserve the activation dtype.""" - output = F.layer_norm(hidden_states.float(), - (self.hidden_size, ), - self.weight, - self.bias, - self.eps) - return output.to(hidden_states.dtype) - - class RMSNorm(nn.Module): """RMS Norm with add residual.""" diff --git a/tests/pytorch/nn/test_fp32_norm.py b/tests/pytorch/nn/test_fp32_norm.py deleted file mode 100644 index 52c745d763..0000000000 --- a/tests/pytorch/nn/test_fp32_norm.py +++ /dev/null @@ -1,29 +0,0 @@ -# Copyright (c) OpenMMLab. All rights reserved. -import pytest -import torch -import torch.nn.functional as F - -from lmdeploy.pytorch.nn.norm import FP32LayerNorm - - -@pytest.mark.parametrize('device', [ - 'cpu', pytest.param('cuda', marks=pytest.mark.skipif( - not torch.cuda.is_available(), reason='requires CUDA')), -]) -@pytest.mark.parametrize('dtype', [torch.float16, torch.bfloat16, torch.float32]) -@pytest.mark.parametrize('bias', [False, True]) -def test_fp32_layer_norm_preserves_parameters_and_output_dtype(dtype, bias, device): - norm = FP32LayerNorm(16, bias=bias, device=device) - generator = torch.Generator().manual_seed(123) - inputs = torch.randn(2, 3, 16, generator=generator).to(device=device, dtype=dtype) - norm.weight.data.copy_(torch.randn(16, generator=generator)) - if bias: - norm.bias.data.copy_(torch.randn(16, generator=generator)) - - expected = F.layer_norm(inputs.float(), (16,), norm.weight, norm.bias, norm.eps).to(dtype) - actual = norm(inputs) - - assert norm.weight.dtype == torch.float32 - assert norm.bias is None or norm.bias.dtype == torch.float32 - assert actual.dtype == dtype - torch.testing.assert_close(actual, expected, rtol=0, atol=0) diff --git a/tests/pytorch/nn/test_glm5_layer_norm.py b/tests/pytorch/nn/test_glm5_layer_norm.py new file mode 100644 index 0000000000..256e4623da --- /dev/null +++ b/tests/pytorch/nn/test_glm5_layer_norm.py @@ -0,0 +1,65 @@ +# Copyright (c) OpenMMLab. All rights reserved. +from types import SimpleNamespace + +import pytest +import torch +import torch.nn.functional as F +from torch import nn + +from lmdeploy.pytorch.models import glm5_next +from lmdeploy.pytorch.nn.kpool import KPOOL_INDEXER_PARAMETER_NAMES, KPoolIndexer +from lmdeploy.pytorch.weight_loader.model_weight_loader import load_weight + + +@pytest.mark.parametrize('device', [ + 'cpu', pytest.param('cuda', marks=pytest.mark.skipif( + not torch.cuda.is_available(), reason='requires CUDA')), +]) +@pytest.mark.parametrize('dtype', [torch.float16, torch.bfloat16, torch.float32]) +@pytest.mark.parametrize('owner', ['kpool', 'vision_merger']) +def test_layer_norm_call_sites_preserve_fp32_contract(monkeypatch, dtype, device, owner): + if owner == 'kpool': + module = KPoolIndexer(128, 2, 128, 8, 4, 4, dtype=dtype, device=device) + assert set(dict(module.named_parameters())) == set(KPOOL_INDEXER_PARAMETER_NAMES) + module.wk = nn.Identity() + norm = module.k_norm + forward = module.project_key + else: + module = glm5_next.Glm5NextVisionPatchMerger( + SimpleNamespace(out_hidden_size=128, intermediate_size=128, swiglu_limit=10.0), + dtype=dtype, device=device) + # Isolate the real merger's norm -> output cast -> GELU call sequence. + module.proj = nn.Identity() + module.gate_up_proj = nn.Identity() + module.down_proj = nn.Identity() + monkeypatch.setattr(glm5_next, '_glm_swiglu_impl', lambda x, *args, **kwargs: x) + norm = module.post_projection_norm + forward = module + + assert type(norm) is nn.LayerNorm + assert norm.eps == 1e-6 + assert norm.weight.dtype == norm.bias.dtype == torch.float32 + assert not norm.weight.requires_grad and not norm.bias.requires_grad + torch.testing.assert_close(norm.weight, torch.ones_like(norm.weight)) + torch.testing.assert_close(norm.bias, torch.zeros_like(norm.bias)) + generator = torch.Generator().manual_seed(123) + # Non-BF16-representable weights catch accidental parameter downcasting. + load_weight(norm.weight, torch.randn(128, generator=generator)) + load_weight(norm.bias, torch.randn(128, generator=generator)) + for tokens in (1, 129, 8192): + inputs = torch.randn(tokens, 128, generator=generator).to(device=device, dtype=dtype) + expected = F.layer_norm(inputs.float(), (128,), norm.weight, norm.bias, norm.eps).to(dtype) + if owner == 'vision_merger': + expected = F.gelu(expected) + actual = forward(inputs) + assert actual.dtype == dtype + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + + +def test_glm_indexer_uses_default_layer_norm(): + config = SimpleNamespace(hidden_size=16, index_n_heads=2, index_head_dim=8, + index_topk=8, q_lora_rank=4, index_kpool=4) + indexer = glm5_next.Glm5NextSparseAttention._build_indexer( + None, config, layer_idx=0, dtype=torch.bfloat16, device=torch.device('cpu')) + assert type(indexer.k_norm) is nn.LayerNorm + assert indexer.k_norm.weight.dtype == indexer.k_norm.bias.dtype == torch.float32 From 17dec3e22d4ce89a31d183d4274e97a3e5bd2026 Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Wed, 23 Sep 2026 12:22:36 +0000 Subject: [PATCH 11/39] refactor: reuse LayerNorm in GLM vision merger --- lmdeploy/pytorch/models/glm5_next.py | 5 +++-- tests/pytorch/nn/test_glm5_layer_norm.py | 20 +++++++++++--------- 2 files changed, 14 insertions(+), 11 deletions(-) diff --git a/lmdeploy/pytorch/models/glm5_next.py b/lmdeploy/pytorch/models/glm5_next.py index c64a6d3798..1cf9c964d2 100644 --- a/lmdeploy/pytorch/models/glm5_next.py +++ b/lmdeploy/pytorch/models/glm5_next.py @@ -37,6 +37,7 @@ HcPrePost, Kda, KPoolIndexer, + LayerNorm, ParallelLMHead, RMSNorm, apply_rotary_pos_emb_fp32, @@ -230,8 +231,8 @@ def __init__(self, bias=False, dtype=dtype, device=device) - self.post_projection_norm = nn.LayerNorm( - dim, eps=1e-6, dtype=torch.float32, device=device).requires_grad_(False) + self.post_projection_norm = LayerNorm(dim, eps=1e-6, dtype=torch.float32, device=device) + nn.init.zeros_(self.post_projection_norm.bias) self.gate_up_proj = build_merged_colwise_linear( in_features=dim, all_out_features=[context_dim, context_dim], diff --git a/tests/pytorch/nn/test_glm5_layer_norm.py b/tests/pytorch/nn/test_glm5_layer_norm.py index 256e4623da..30f9c64f02 100644 --- a/tests/pytorch/nn/test_glm5_layer_norm.py +++ b/tests/pytorch/nn/test_glm5_layer_norm.py @@ -7,6 +7,7 @@ from torch import nn from lmdeploy.pytorch.models import glm5_next +from lmdeploy.pytorch.nn import LayerNorm from lmdeploy.pytorch.nn.kpool import KPOOL_INDEXER_PARAMETER_NAMES, KPoolIndexer from lmdeploy.pytorch.weight_loader.model_weight_loader import load_weight @@ -16,8 +17,8 @@ not torch.cuda.is_available(), reason='requires CUDA')), ]) @pytest.mark.parametrize('dtype', [torch.float16, torch.bfloat16, torch.float32]) -@pytest.mark.parametrize('owner', ['kpool', 'vision_merger']) -def test_layer_norm_call_sites_preserve_fp32_contract(monkeypatch, dtype, device, owner): +@pytest.mark.parametrize(('owner', 'hidden_size'), [('kpool', 128), ('vision_merger', 128), ('vision_merger', 4096)]) +def test_layer_norm_call_sites_preserve_fp32_contract(monkeypatch, dtype, device, owner, hidden_size): if owner == 'kpool': module = KPoolIndexer(128, 2, 128, 8, 4, 4, dtype=dtype, device=device) assert set(dict(module.named_parameters())) == set(KPOOL_INDEXER_PARAMETER_NAMES) @@ -26,7 +27,7 @@ def test_layer_norm_call_sites_preserve_fp32_contract(monkeypatch, dtype, device forward = module.project_key else: module = glm5_next.Glm5NextVisionPatchMerger( - SimpleNamespace(out_hidden_size=128, intermediate_size=128, swiglu_limit=10.0), + SimpleNamespace(out_hidden_size=hidden_size, intermediate_size=128, swiglu_limit=10.0), dtype=dtype, device=device) # Isolate the real merger's norm -> output cast -> GELU call sequence. module.proj = nn.Identity() @@ -36,19 +37,20 @@ def test_layer_norm_call_sites_preserve_fp32_contract(monkeypatch, dtype, device norm = module.post_projection_norm forward = module - assert type(norm) is nn.LayerNorm - assert norm.eps == 1e-6 + assert type(norm) is (nn.LayerNorm if owner == 'kpool' else LayerNorm) + assert (norm.eps if owner == 'kpool' else norm.impl.eps) == 1e-6 + assert set(norm.state_dict()) == {'weight', 'bias'} assert norm.weight.dtype == norm.bias.dtype == torch.float32 assert not norm.weight.requires_grad and not norm.bias.requires_grad torch.testing.assert_close(norm.weight, torch.ones_like(norm.weight)) torch.testing.assert_close(norm.bias, torch.zeros_like(norm.bias)) generator = torch.Generator().manual_seed(123) # Non-BF16-representable weights catch accidental parameter downcasting. - load_weight(norm.weight, torch.randn(128, generator=generator)) - load_weight(norm.bias, torch.randn(128, generator=generator)) + load_weight(norm.weight, torch.randn(hidden_size, generator=generator)) + load_weight(norm.bias, torch.randn(hidden_size, generator=generator)) for tokens in (1, 129, 8192): - inputs = torch.randn(tokens, 128, generator=generator).to(device=device, dtype=dtype) - expected = F.layer_norm(inputs.float(), (128,), norm.weight, norm.bias, norm.eps).to(dtype) + inputs = torch.randn(tokens, hidden_size, generator=generator).to(device=device, dtype=dtype) + expected = F.layer_norm(inputs.float(), (hidden_size,), norm.weight, norm.bias, 1e-6).to(dtype) if owner == 'vision_merger': expected = F.gelu(expected) actual = forward(inputs) From 7b5a819834032a2a6657111469776f5d5957bc4a Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Wed, 23 Sep 2026 12:31:38 +0000 Subject: [PATCH 12/39] refactor: reuse common FP32 rotary operator for GLM vision Replace the compiled GLM helper with an opt-in FP32 compute contract through ApplyRotaryEmb and its backend build spec. CUDA retains FP32 arithmetic inside the fused kernel; defaults remain unchanged and unsupported Dlinfer requests fail explicitly. Validated 73 focused tests and actual vLLM rotary kernel parity with identical inputs and tables. FP32-table outputs can differ from the old compiled helper at rounding boundaries. --- lmdeploy/pytorch/backends/apply_rotary_emb.py | 2 + .../pytorch/backends/cuda/apply_rotary_emb.py | 5 +- lmdeploy/pytorch/backends/cuda/op_backend.py | 2 +- .../backends/default/apply_rotary_emb.py | 11 ++++ .../pytorch/backends/default/op_backend.py | 2 +- .../pytorch/backends/dlinfer/op_backend.py | 2 + .../kernels/cuda/apply_rotary_pos_emb.py | 23 +++++-- lmdeploy/pytorch/models/glm5_next.py | 5 +- lmdeploy/pytorch/nn/__init__.py | 1 - lmdeploy/pytorch/nn/rotary_embedding.py | 35 ++-------- tests/pytorch/nn/test_rotary_embedding.py | 64 +++++++++++++++---- 11 files changed, 96 insertions(+), 56 deletions(-) diff --git a/lmdeploy/pytorch/backends/apply_rotary_emb.py b/lmdeploy/pytorch/backends/apply_rotary_emb.py index cdda460055..76c9c2cdde 100644 --- a/lmdeploy/pytorch/backends/apply_rotary_emb.py +++ b/lmdeploy/pytorch/backends/apply_rotary_emb.py @@ -25,3 +25,5 @@ def forward(self, @dataclass(frozen=True) class ApplyRotaryEmbBuildSpec(BuildSpec[ApplyRotaryEmbImpl]): """Request construction of an apply-RoPE operator.""" + + enable_fp32_compute: bool = False diff --git a/lmdeploy/pytorch/backends/cuda/apply_rotary_emb.py b/lmdeploy/pytorch/backends/cuda/apply_rotary_emb.py index f73b3d158e..36ee5d8d49 100644 --- a/lmdeploy/pytorch/backends/cuda/apply_rotary_emb.py +++ b/lmdeploy/pytorch/backends/cuda/apply_rotary_emb.py @@ -10,6 +10,9 @@ class TritonApplyRotaryEmbImpl(ApplyRotaryEmbImpl): """Apply rotary embedding implementation.""" + def __init__(self, enable_fp32_compute: bool = False): + self.enable_fp32_compute = enable_fp32_compute + def forward(self, query: Tensor, key: Tensor, @@ -25,4 +28,4 @@ def forward(self, q_embed = torch.empty_like(query) k_embed = torch.empty_like(key) return apply_rotary_pos_emb(query, key, cos, sin, q_embed, k_embed, - complex_mode=complex_mode) + complex_mode=complex_mode, enable_fp32_compute=self.enable_fp32_compute) diff --git a/lmdeploy/pytorch/backends/cuda/op_backend.py b/lmdeploy/pytorch/backends/cuda/op_backend.py index 2105c61a5c..3294b54f4c 100644 --- a/lmdeploy/pytorch/backends/cuda/op_backend.py +++ b/lmdeploy/pytorch/backends/cuda/op_backend.py @@ -68,7 +68,7 @@ def build_op(cls, spec: BuildSpec[ImplT], *, enable_deterministic: bool = False) return cast(ImplT, CudaKdaImpl()) if isinstance(spec, ApplyRotaryEmbBuildSpec): from .apply_rotary_emb import TritonApplyRotaryEmbImpl - return cast(ImplT, TritonApplyRotaryEmbImpl()) + return cast(ImplT, TritonApplyRotaryEmbImpl(enable_fp32_compute=spec.enable_fp32_compute)) if isinstance(spec, RMSNormBuildSpec): from .norm import TritonRMSNormImpl return cast(ImplT, TritonRMSNormImpl(spec.hidden_size, spec.eps)) diff --git a/lmdeploy/pytorch/backends/default/apply_rotary_emb.py b/lmdeploy/pytorch/backends/default/apply_rotary_emb.py index 138c64c463..cfcbc0b302 100644 --- a/lmdeploy/pytorch/backends/default/apply_rotary_emb.py +++ b/lmdeploy/pytorch/backends/default/apply_rotary_emb.py @@ -48,6 +48,9 @@ def _prepare_cos_sin(query: Tensor, cos: Tensor, sin: Tensor, complex_mode: bool class DefaultApplyRotaryEmbImpl(ApplyRotaryEmbImpl): """Apply rotary embedding implementation.""" + def __init__(self, enable_fp32_compute: bool = False): + self.enable_fp32_compute = enable_fp32_compute + def forward(self, query: Tensor, key: Tensor, @@ -61,6 +64,10 @@ def forward(self, else: rotate_fn = rotate_half cos, sin = _prepare_cos_sin(query, cos, sin, complex_mode) + original_query, original_key = query, key + if self.enable_fp32_compute: + query, key = query.float(), key.float() + cos, sin = cos.float(), sin.float() if inplace: q_embed = query k_embed = key @@ -73,4 +80,8 @@ def forward(self, else: q_embed = (query * cos) + (rotate_fn(query) * sin) k_embed = (key * cos) + (rotate_fn(key) * sin) + if self.enable_fp32_compute: + q_embed, k_embed = q_embed.to(original_query.dtype), k_embed.to(original_key.dtype) + if inplace: + q_embed, k_embed = original_query.copy_(q_embed), original_key.copy_(k_embed) return q_embed, k_embed diff --git a/lmdeploy/pytorch/backends/default/op_backend.py b/lmdeploy/pytorch/backends/default/op_backend.py index 8e8f5c57ee..408062ce35 100644 --- a/lmdeploy/pytorch/backends/default/op_backend.py +++ b/lmdeploy/pytorch/backends/default/op_backend.py @@ -38,7 +38,7 @@ def build_op(cls, spec: BuildSpec[ImplT], *, enable_deterministic: bool = False) return cast(ImplT, _build_rotary_embedding(spec)) if isinstance(spec, ApplyRotaryEmbBuildSpec): from .apply_rotary_emb import DefaultApplyRotaryEmbImpl - return cast(ImplT, DefaultApplyRotaryEmbImpl()) + return cast(ImplT, DefaultApplyRotaryEmbImpl(enable_fp32_compute=spec.enable_fp32_compute)) if isinstance(spec, RMSNormBuildSpec): from .norm import DefaultRMSNormImpl return cast(ImplT, DefaultRMSNormImpl(spec.hidden_size, spec.eps)) diff --git a/lmdeploy/pytorch/backends/dlinfer/op_backend.py b/lmdeploy/pytorch/backends/dlinfer/op_backend.py index 7b4b838526..44582593be 100644 --- a/lmdeploy/pytorch/backends/dlinfer/op_backend.py +++ b/lmdeploy/pytorch/backends/dlinfer/op_backend.py @@ -38,6 +38,8 @@ def build_op(cls, spec: BuildSpec[ImplT], *, enable_deterministic: bool = False) from .activation import DlinferSiluAndMulImpl return cast(ImplT, DlinferSiluAndMulImpl()) if isinstance(spec, ApplyRotaryEmbBuildSpec): + if spec.enable_fp32_compute: + raise NotImplementedError('Dlinfer ApplyRotaryEmb does not support enable_fp32_compute=True.') from .apply_rotary_emb import DlinferApplyRotaryEmbImpl return cast(ImplT, DlinferApplyRotaryEmbImpl()) if isinstance(spec, RMSNormBuildSpec): diff --git a/lmdeploy/pytorch/kernels/cuda/apply_rotary_pos_emb.py b/lmdeploy/pytorch/kernels/cuda/apply_rotary_pos_emb.py index 4ec9b29ab7..32de05d4a0 100644 --- a/lmdeploy/pytorch/kernels/cuda/apply_rotary_pos_emb.py +++ b/lmdeploy/pytorch/kernels/cuda/apply_rotary_pos_emb.py @@ -6,11 +6,15 @@ @triton.jit -def _apply_rotary_impl(x_l, x_h, cos_l, cos_h, sin_l, sin_h): +def _apply_rotary_impl(x_l, x_h, cos_l, cos_h, sin_l, sin_h, ENABLE_FP32_COMPUTE: tl.constexpr = False): """Apply rotary positional embedding implementation.""" # x_l, x_h: [BLOCK, BLOCK_N] # cos_l, cos_h, sin_l, sin_h: [BLOCK, BLOCK_N] + if ENABLE_FP32_COMPUTE: + # Match FlashAttention RoPE's FP32 multiply-add path. + return x_l * cos_l - x_h * sin_l, x_h * cos_h + x_l * sin_h + # triton 3.4 would do fma 3 times to perform the above computation, # which causes higher numerical error. So we manually expand the # computation to avoid fma. @@ -47,6 +51,7 @@ def apply_rotary_pos_emb_qk_kernel( BLOCK_QH: tl.constexpr, BLOCK_N: tl.constexpr, COMPLEX: tl.constexpr = False, + ENABLE_FP32_COMPUTE: tl.constexpr = False, ): """Apply rotary on key AND query kernel.""" seq_block_id = tl.program_id(1) @@ -77,7 +82,7 @@ def apply_rotary_pos_emb_qk_kernel( seq_mask = pos_mask[:, None] & feat_mask[None, :] cs_offset_l = pos_offset[:, None] * cs_stride + feat_offset_l[None, :] cs_offset_h = pos_offset[:, None] * cs_stride + feat_offset_h[None, :] - q_elem_type = Q.dtype.element_ty + q_elem_type = tl.float32 if ENABLE_FP32_COMPUTE else Q.dtype.element_ty cos_l = tl.load(COS + cs_offset_l).to(q_elem_type) cos_h = tl.load(COS + cs_offset_h).to(q_elem_type) sin_l = tl.load(SIN + cs_offset_l).to(q_elem_type) @@ -97,8 +102,10 @@ def apply_rotary_pos_emb_qk_kernel( q_l = tl.load(ql_ptrs) q_h = tl.load(qh_ptrs) + if ENABLE_FP32_COMPUTE: + q_l, q_h = q_l.to(tl.float32), q_h.to(tl.float32) - qe_l, qe_h = _apply_rotary_impl(q_l, q_h, cos_l, cos_h, sin_l, sin_h) + qe_l, qe_h = _apply_rotary_impl(q_l, q_h, cos_l, cos_h, sin_l, sin_h, ENABLE_FP32_COMPUTE) tl.store(qel_ptrs, qe_l, mask=seq_mask) tl.store(qeh_ptrs, qe_h, mask=seq_mask) @@ -116,8 +123,10 @@ def apply_rotary_pos_emb_qk_kernel( keh_ptrs += head_id * stride_keh k_l = tl.load(kl_ptrs) k_h = tl.load(kh_ptrs) + if ENABLE_FP32_COMPUTE: + k_l, k_h = k_l.to(tl.float32), k_h.to(tl.float32) - ke_l, ke_h = _apply_rotary_impl(k_l, k_h, cos_l, cos_h, sin_l, sin_h) + ke_l, ke_h = _apply_rotary_impl(k_l, k_h, cos_l, cos_h, sin_l, sin_h, ENABLE_FP32_COMPUTE) tl.store(kel_ptrs, ke_l, mask=seq_mask) tl.store(keh_ptrs, ke_h, mask=seq_mask) @@ -129,7 +138,8 @@ def apply_rotary_pos_emb(q: Tensor, sin: Tensor, q_embed: Tensor = None, k_embed: Tensor = None, - complex_mode: bool = False): + complex_mode: bool = False, + enable_fp32_compute: bool = False): """Apply rotary positional embedding on query and key. Args: @@ -144,6 +154,8 @@ def apply_rotary_pos_emb(q: Tensor, cos/sin should be (seq_len, dim//2). If False (default), use rotate_half style where front/back halves are paired. cos/sin should be (seq_len, dim). + enable_fp32_compute (bool): Compute Q/K and table products in FP32, + then cast on store. Defaults to False. Returns: tuple[Tensor, Tensor]: Embedded query and key. @@ -215,6 +227,7 @@ def apply_rotary_pos_emb(q: Tensor, BLOCK_QH=num_heads_q, BLOCK_N=BLOCK_N, COMPLEX=complex_mode, + ENABLE_FP32_COMPUTE=enable_fp32_compute, num_warps=num_warps, num_stages=num_stages) diff --git a/lmdeploy/pytorch/models/glm5_next.py b/lmdeploy/pytorch/models/glm5_next.py index 1cf9c964d2..9ab81f1ac2 100644 --- a/lmdeploy/pytorch/models/glm5_next.py +++ b/lmdeploy/pytorch/models/glm5_next.py @@ -33,6 +33,7 @@ from lmdeploy.pytorch.engine.cache_engine.schema import BlockCacheRequest from lmdeploy.pytorch.model_inputs import StepContext, StepContextManager, get_step_ctx_manager from lmdeploy.pytorch.nn import ( + ApplyRotaryEmb, FlashAttention, HcPrePost, Kda, @@ -40,7 +41,6 @@ LayerNorm, ParallelLMHead, RMSNorm, - apply_rotary_pos_emb_fp32, ) from lmdeploy.pytorch.nn.gated_delta import GatedDeltaMeta, GatedDeltaMetaBuilder, build_rmsnorm_gated from lmdeploy.pytorch.nn.kpool import ( @@ -140,6 +140,7 @@ def __init__(self, # Reuse LMDeploy's QKV/row-parallel projections and rotary operator. # Vision weights stay in BF16 even when the language tower is block-FP8. super().__init__(config, dtype=dtype, device=device) + self.apply_rotary_pos_emb = ApplyRotaryEmb(enable_fp32_compute=True) self.q_norm = Glm5NextVisionRMSNorm(self.head_dim, eps=self.qk_norm_eps, quant_config=None, @@ -165,7 +166,7 @@ def forward( query = self.q_norm(query) key = self.k_norm(key) cos, sin = rotary_pos_emb - query, key = apply_rotary_pos_emb_fp32(query, key, cos, sin) + query, key = self.apply_rotary_pos_emb(query, key, cos, sin, inplace=False) output = self.attention( query, key, diff --git a/lmdeploy/pytorch/nn/__init__.py b/lmdeploy/pytorch/nn/__init__.py index f31de8f881..1feb2c55e2 100644 --- a/lmdeploy/pytorch/nn/__init__.py +++ b/lmdeploy/pytorch/nn/__init__.py @@ -12,7 +12,6 @@ ApplyRotaryEmb, # noqa: F401 RopeType, # noqa: F401 YarnParameters, # noqa: F401 - apply_rotary_pos_emb_fp32, # noqa: F401 build_rotary_embedding, # noqa: F401 build_rotary_embedding_from_config, # noqa: F401 build_rotary_params, # noqa: F401 diff --git a/lmdeploy/pytorch/nn/rotary_embedding.py b/lmdeploy/pytorch/nn/rotary_embedding.py index e94ce43de3..2757e2c085 100644 --- a/lmdeploy/pytorch/nn/rotary_embedding.py +++ b/lmdeploy/pytorch/nn/rotary_embedding.py @@ -244,41 +244,14 @@ def build_rotary_embedding_from_config(config: PretrainedConfig, device: torch.d return build_rotary_embedding(**rope_params, device=device) -def _rotate_half(x: Tensor) -> Tensor: - """Rotate the two contiguous halves used by NeoX-style RoPE.""" - x1, x2 = x.chunk(2, dim=-1) - return torch.cat((-x2, x1), dim=-1) - - -@torch.compile(dynamic=True) -def apply_rotary_pos_emb_fp32(query: Tensor, - key: Tensor, - cos: Tensor, - sin: Tensor, - unsqueeze_dim: int = 1) -> tuple[Tensor, Tensor]: - """Apply NeoX-style RoPE with FP32 arithmetic and dtype-preserving output. - - Used by GLM-5.3 vision attention to match its FP32 rotary arithmetic; the text attention uses NoPE and does not call - this helper. - """ - query_dtype = query.dtype - key_dtype = key.dtype - query = query.float() - key = key.float() - cos = cos.unsqueeze(unsqueeze_dim).float() - sin = sin.unsqueeze(unsqueeze_dim).float() - query = query * cos + _rotate_half(query) * sin - key = key * cos + _rotate_half(key) * sin - return query.to(query_dtype), key.to(key_dtype) - - class ApplyRotaryEmb(nn.Module): - """Apply rotary embedding.""" + """Apply rotary embedding, optionally computing in FP32 before the output + cast.""" - def __init__(self): + def __init__(self, enable_fp32_compute: bool = False): super().__init__() self.impl = get_backend().build_op( - ApplyRotaryEmbBuildSpec(), + ApplyRotaryEmbBuildSpec(enable_fp32_compute=enable_fp32_compute), enable_deterministic=get_build_model_context().enable_deterministic, ) diff --git a/tests/pytorch/nn/test_rotary_embedding.py b/tests/pytorch/nn/test_rotary_embedding.py index cea8bb3892..9ecaabde31 100644 --- a/tests/pytorch/nn/test_rotary_embedding.py +++ b/tests/pytorch/nn/test_rotary_embedding.py @@ -234,26 +234,62 @@ def test_default_apply_rotary_complex_accepts_half_width_tables_with_empty_key() not torch.cuda.is_available(), reason='requires CUDA')), ]) @pytest.mark.parametrize('dtype', [torch.float16, torch.bfloat16, torch.float32]) -@pytest.mark.parametrize('unsqueeze_dim', [0, 1]) -def test_fp32_rotary_matches_reference_and_preserves_input(dtype, unsqueeze_dim, device): - from lmdeploy.pytorch.nn.rotary_embedding import apply_rotary_pos_emb_fp32 +@pytest.mark.parametrize('enable_fp32_compute', [False, True]) +@pytest.mark.parametrize('inplace', [False, True]) +@pytest.mark.parametrize('complex_mode', [False, True]) +def test_apply_rotary_compute_precision(monkeypatch, dtype, device, enable_fp32_compute, inplace, complex_mode): + from lmdeploy.pytorch.backends.default.op_backend import DefaultOpsBackend + from lmdeploy.pytorch.nn import rotary_embedding + + if device == 'cpu': + monkeypatch.setattr(rotary_embedding, 'get_backend', lambda: DefaultOpsBackend) + module = rotary_embedding.ApplyRotaryEmb(enable_fp32_compute=enable_fp32_compute) generator = torch.Generator().manual_seed(123) - shape = (3, 5, 16) if unsqueeze_dim == 0 else (5, 3, 16) - query = torch.randn(shape, generator=generator).to(device=device, dtype=dtype) - key = torch.randn(shape, generator=generator).to(device=device, dtype=dtype) - # Exactly representable coefficients isolate intermediate rounding in BF16/FP16. - cos = torch.full((5, 16), 0.625, dtype=dtype, device=device) - sin = torch.full((5, 16), 0.375, dtype=dtype, device=device) + # Unequal head counts and strided Q/K exercise the fused CUDA kernel. + query = torch.randn(33, 3, 32, generator=generator).to(device=device, dtype=dtype)[..., ::2] + key = torch.randn(33, 2, 32, generator=generator).to(device=device, dtype=dtype)[..., ::2] + # FP32 tables also catch accidental downcasting before FP32 arithmetic. + cos_value = 0.625 + (2**-12 if enable_fp32_compute else 0) + sin_value = 0.375 + (2**-13 if enable_fp32_compute else 0) + table_dtype = torch.float32 if enable_fp32_compute else dtype + table_dim = 8 if complex_mode else 16 + cos = torch.full((33, table_dim), cos_value, dtype=table_dtype, device=device) + sin = torch.full((33, table_dim), sin_value, dtype=table_dtype, device=device) original = (query.clone(), key.clone()) - outputs = apply_rotary_pos_emb_fp32(query, key, cos, sin, unsqueeze_dim) + outputs = module(query, key, cos, sin, inplace=inplace, complex_mode=complex_mode) for value, saved, actual in zip((query, key), original, outputs): - left, right = value.float().chunk(2, dim=-1) - expected = torch.cat((left * 0.625 - right * 0.375, - right * 0.625 + left * 0.375), dim=-1).to(dtype) + inputs = saved.float() if enable_fp32_compute else saved + if complex_mode: + rotated = _rotate_complex(inputs) + else: + left, right = inputs.chunk(2, dim=-1) + rotated = torch.cat((-right, left), dim=-1) + expected = (inputs * cos_value + rotated * sin_value).to(dtype) assert actual.dtype == dtype torch.testing.assert_close(actual, expected, rtol=1e-6 if dtype == torch.float32 else 0, atol=1e-7 if dtype == torch.float32 else 0) - torch.testing.assert_close(value, saved, rtol=0, atol=0) + if inplace: + assert actual is value + else: + torch.testing.assert_close(value, saved, rtol=0, atol=0) + + +def test_dlinfer_rejects_fp32_rotary(): + from lmdeploy.pytorch.backends.apply_rotary_emb import ApplyRotaryEmbBuildSpec + from lmdeploy.pytorch.backends.dlinfer.op_backend import DlinferOpsBackend + + with pytest.raises(NotImplementedError, match='enable_fp32_compute=True'): + DlinferOpsBackend.build_op(ApplyRotaryEmbBuildSpec(enable_fp32_compute=True)) + + +def test_glm_vision_uses_common_fp32_rotary(): + from lmdeploy.pytorch.models.glm5_next import Glm5NextVisionAttention + from lmdeploy.pytorch.nn import ApplyRotaryEmb + + config = PretrainedConfig(hidden_size=128, num_heads=2, attention_bias=False) + module = Glm5NextVisionAttention(config, dtype=torch.bfloat16, device='cpu') + assert type(module.apply_rotary_pos_emb) is ApplyRotaryEmb + assert module.apply_rotary_pos_emb.impl.enable_fp32_compute From c1ae34b3742af59436a1d49bc743ec67e6b5bcea Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Wed, 23 Sep 2026 13:07:19 +0000 Subject: [PATCH 13/39] perf: preserve strided KDA decode inputs Reuse the existing strided TileLang recurrent inputs instead of copying Q/K/V in AR and speculative decode. Preserve contiguous prefill inputs and explicitly allocate contiguous recurrent output for channel-major views. Validated 25 new stride/state/graph tests, 25 shared GDR tests, and 12 complete KDA adapter comparisons. TP4 MTP-off text regression matches the previous head exactly at 124 full-vocabulary logits positions and all 29 requests. --- lmdeploy/pytorch/backends/cuda/kda.py | 16 ++--- .../pytorch/kernels/cuda/gated_delta_rule.py | 3 +- tests/pytorch/kernel/test_kda_strided.py | 67 +++++++++++++++++++ 3 files changed, 77 insertions(+), 9 deletions(-) create mode 100644 tests/pytorch/kernel/test_kda_strided.py diff --git a/lmdeploy/pytorch/backends/cuda/kda.py b/lmdeploy/pytorch/backends/cuda/kda.py index f7603c11b4..202a6bde12 100644 --- a/lmdeploy/pytorch/backends/cuda/kda.py +++ b/lmdeploy/pytorch/backends/cuda/kda.py @@ -141,7 +141,7 @@ def _forward_spec_decode(self, mixed_qkv, raw_gate, raw_beta, conv_state, cache_seqlens=history) mixed = mixed.transpose(1, 2) heads, dim = kwargs['num_heads'], kwargs['head_dim'] - q, k, v = [x.reshape(batch, steps, heads, dim).contiguous() + q, k, v = [x.reshape(batch, steps, heads, dim) for x in mixed.split(heads * dim, dim=-1)] gate = self.kda_gate(raw_gate.reshape(batch, steps, heads, dim).contiguous(), kwargs['a_log'], kwargs['dt_bias'], lower_bound=kwargs['lower_bound']) @@ -221,9 +221,9 @@ def forward( mixed_qkv = self._conv(mixed_qkv, conv_weight, conv_bias, conv_state, metadata) q, k, v = mixed_qkv.split(num_heads * head_dim, dim=-1) - q = q.unflatten(-1, (num_heads, head_dim)).contiguous() - k = k.unflatten(-1, (num_heads, head_dim)).contiguous() - v = v.unflatten(-1, (num_heads, head_dim)).contiguous() + q = q.unflatten(-1, (num_heads, head_dim)) + k = k.unflatten(-1, (num_heads, head_dim)) + v = v.unflatten(-1, (num_heads, head_dim)) raw_gate = raw_gate.unflatten( -1, (num_heads, head_dim)).contiguous() raw_beta = raw_beta.contiguous() @@ -232,7 +232,7 @@ def forward( if metadata.is_decoding: def decode_view(x: torch.Tensor) -> torch.Tensor: return x.squeeze(0).unflatten( - 0, (batch_size, 1)).contiguous() + 0, (batch_size, 1)) output, final_state = self.fused_recurrent_kda( q=decode_view(q), @@ -253,9 +253,9 @@ def decode_view(x: torch.Tensor) -> torch.Tensor: output = output.flatten(0, 1).unsqueeze(0) else: output, final_state = self.chunk_kda( - q=q, - k=k, - v=v, + q=q.contiguous(), + k=k.contiguous(), + v=v.contiguous(), g=raw_gate, beta=raw_beta, A_log=a_log, diff --git a/lmdeploy/pytorch/kernels/cuda/gated_delta_rule.py b/lmdeploy/pytorch/kernels/cuda/gated_delta_rule.py index 481829bd7d..c416142507 100644 --- a/lmdeploy/pytorch/kernels/cuda/gated_delta_rule.py +++ b/lmdeploy/pytorch/kernels/cuda/gated_delta_rule.py @@ -663,7 +663,8 @@ def fused_recurrent_gated_delta_rule( 'cache_seqlens must be on the same device as q' assert initial_state is not None, 'initial_state is required' - o = torch.empty_like(v) + # Strided inputs may be dense transposes, but Out uses a contiguous Tensor. + o = torch.empty_like(v, memory_format=torch.contiguous_format) final_state = initial_state state_dtype = q.dtype if final_state is not None: diff --git a/tests/pytorch/kernel/test_kda_strided.py b/tests/pytorch/kernel/test_kda_strided.py new file mode 100644 index 0000000000..5889fec4d8 --- /dev/null +++ b/tests/pytorch/kernel/test_kda_strided.py @@ -0,0 +1,67 @@ +# Copyright (c) OpenMMLab. All rights reserved. +import pytest +import torch + + +@pytest.mark.parametrize('batch', [1, 3]) +@pytest.mark.parametrize('steps', [1, 3, 6]) +@pytest.mark.parametrize('channel_major', [False, True]) +@pytest.mark.parametrize('state_dtype', [torch.float32, torch.bfloat16]) +def test_kda_strided_inputs_match_contiguous_state_ring(batch, steps, channel_major, state_dtype): + from lmdeploy.pytorch.kernels.cuda.gated_delta_rule import fused_recurrent_gated_delta_rule + + torch.manual_seed(17) + heads, dim, ring = 2, 128, 6 + if channel_major: + mixed = torch.randn(batch, 3 * heads * dim, steps, device='cuda', dtype=torch.bfloat16).transpose(1, 2) + else: + mixed = torch.randn(batch, steps, 3 * heads * dim, device='cuda', dtype=torch.bfloat16) + q, k, v = [x.unflatten(-1, (heads, dim)) for x in mixed.chunk(3, dim=-1)] + gate = -torch.rand(batch, steps, heads, dim, device='cuda') + beta = torch.rand(batch, steps, heads, device='cuda') + initial = torch.randn(3, ring, heads, dim, dim, device='cuda', dtype=state_dtype) * 0.1 + ids = torch.tensor([2, -1, 0][:batch], device='cuda') + history = torch.tensor([5, 0, 3][:batch], device='cuda', dtype=torch.int32) + kwargs = dict(g=gate, beta=beta, state_indices=ids, cache_seqlens=history, + output_final_state=True, transpose_state_layout=True, use_qk_l2norm_in_kernel=True) + expected, expected_state = fused_recurrent_gated_delta_rule( + q.contiguous(), k.contiguous(), v.contiguous(), initial_state=initial.clone(), **kwargs) + actual, actual_state = fused_recurrent_gated_delta_rule(q, k, v, initial_state=initial.clone(), **kwargs) + assert actual.is_contiguous() + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + torch.testing.assert_close(actual_state, expected_state, rtol=0, atol=0) + torch.testing.assert_close(actual_state[1], initial[1], rtol=0, atol=0) + if batch == 3: + assert torch.count_nonzero(actual[1]) == 0 + + +def test_kda_strided_graph_replay_changes_inputs_and_history(): + from lmdeploy.pytorch.kernels.cuda.gated_delta_rule import fused_recurrent_gated_delta_rule + + torch.manual_seed(31) + heads, dim, steps = 2, 128, 6 + mixed = torch.randn(1, 3 * heads * dim, steps, device='cuda', dtype=torch.bfloat16) + q, k, v = [x.unflatten(-1, (heads, dim)) for x in mixed.transpose(1, 2).chunk(3, dim=-1)] + gate = -torch.rand_like(q, dtype=torch.float32).contiguous() + beta = torch.rand(1, steps, heads, device='cuda') + initial = torch.randn(1, steps, heads, dim, dim, device='cuda') * 0.1 + state = initial.clone() + history = torch.tensor([5], device='cuda', dtype=torch.int32) + kwargs = dict(g=gate, beta=beta, cache_seqlens=history, output_final_state=True, + transpose_state_layout=True, use_qk_l2norm_in_kernel=True) + + def run(): + return fused_recurrent_gated_delta_rule(q, k, v, initial_state=state, **kwargs) + + run() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + actual, _ = run() + mixed.normal_() + history.fill_(2) + state.copy_(initial) + expected, expected_state = fused_recurrent_gated_delta_rule( + q.contiguous(), k.contiguous(), v.contiguous(), initial_state=initial.clone(), **kwargs) + graph.replay() + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + torch.testing.assert_close(state, expected_state, rtol=0, atol=0) From 7fa48ebb14f066619ffd9756ceb7bf65d89284dd Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Wed, 23 Sep 2026 13:38:45 +0000 Subject: [PATCH 14/39] perf: fuse KDA gates while preserving sigmoid precision --- lmdeploy/pytorch/backends/cuda/kda.py | 27 ++-- .../pytorch/kernels/cuda/gated_delta_rule.py | 59 ++++++++- tests/pytorch/kernel/test_kda_fused_gate.py | 120 ++++++++++++++++++ 3 files changed, 192 insertions(+), 14 deletions(-) create mode 100644 tests/pytorch/kernel/test_kda_fused_gate.py diff --git a/lmdeploy/pytorch/backends/cuda/kda.py b/lmdeploy/pytorch/backends/cuda/kda.py index 202a6bde12..7c94a4b07e 100644 --- a/lmdeploy/pytorch/backends/cuda/kda.py +++ b/lmdeploy/pytorch/backends/cuda/kda.py @@ -106,12 +106,18 @@ def store(state, values, lengths): return output def _decode_recurrent(self, q, k, v, g, beta, A_log, dt_bias, initial_state, - output_final_state=True, lower_bound=None, **kwargs): + output_final_state=True, lower_bound=None, state_indices=None, cache_seqlens=None, **kwargs): """Keep AR and MTP on the same recurrence and gate arithmetic.""" - gate = self.kda_gate(g, A_log, dt_bias, lower_bound=lower_bound) - return self.recurrent_func(q, k, v, g=gate, beta=beta.float().sigmoid(), + gate_args = {} + if lower_bound is not None: + gate_args = dict(a_log=A_log, dt_bias=dt_bias, lower_bound=lower_bound) + else: + g = self.kda_gate(g, A_log, dt_bias) + beta = beta.float().sigmoid() + return self.recurrent_func(q, k, v, g=g, beta=beta, initial_state=initial_state, output_final_state=output_final_state, - use_qk_l2norm_in_kernel=True, transpose_state_layout=True) + use_qk_l2norm_in_kernel=True, transpose_state_layout=True, + state_indices=state_indices, cache_seqlens=cache_seqlens, **gate_args) def _forward_spec_decode(self, mixed_qkv, raw_gate, raw_beta, conv_state, recurrent_state, metadata, **kwargs): @@ -143,13 +149,12 @@ def _forward_spec_decode(self, mixed_qkv, raw_gate, raw_beta, conv_state, heads, dim = kwargs['num_heads'], kwargs['head_dim'] q, k, v = [x.reshape(batch, steps, heads, dim) for x in mixed.split(heads * dim, dim=-1)] - gate = self.kda_gate(raw_gate.reshape(batch, steps, heads, dim).contiguous(), - kwargs['a_log'], kwargs['dt_bias'], lower_bound=kwargs['lower_bound']) - beta = raw_beta.reshape(batch, steps, heads).float().sigmoid() - output, _ = self.recurrent_func(q, k, v, g=gate, beta=beta, - initial_state=recurrent_state, state_indices=signed_ids, - cache_seqlens=history, output_final_state=True, - use_qk_l2norm_in_kernel=True, transpose_state_layout=True) + output, _ = self._decode_recurrent( + q, k, v, g=raw_gate.reshape(batch, steps, heads, dim).contiguous(), + beta=raw_beta.reshape(batch, steps, heads).contiguous(), + A_log=kwargs['a_log'], dt_bias=kwargs['dt_bias'], lower_bound=kwargs['lower_bound'], + initial_state=recurrent_state, state_indices=signed_ids, cache_seqlens=history, + output_final_state=True) return output.reshape(1, batch * steps, heads, dim) def _conv( diff --git a/lmdeploy/pytorch/kernels/cuda/gated_delta_rule.py b/lmdeploy/pytorch/kernels/cuda/gated_delta_rule.py index c416142507..c34596f806 100644 --- a/lmdeploy/pytorch/kernels/cuda/gated_delta_rule.py +++ b/lmdeploy/pytorch/kernels/cuda/gated_delta_rule.py @@ -279,7 +279,10 @@ def fused_recurrent_gated_delta_rule_fwd(SEQLEN, is_circular_buffer: bool = False, transpose_state_layout: bool = False, num_warps: int = 1, - channelwise_g: bool = False): + channelwise_g: bool = False, + fuse_kda_gate: bool = False, + kda_lower_bound: float = -5.0, + has_dt_bias: bool = False): """Build the layout-specific recurrent GDR TileLang kernel. Common compile-time metadata is computed once here. The only structural branch is the returned T.prim_func body, @@ -466,8 +469,12 @@ def fused_recurrent_gated_delta_rule_transposed_main( State: T.StridedTensor([N, NUM_STATE, HV, V, K], dtype=state_dtype, strides=state_stride), StateIndices: T.Tensor([B], dtype=torch.int64) = None, CacheSeqlens: T.Tensor([B], dtype=torch.int32) = None, + ALog: T.Tensor([HV], dtype=torch.float32) = None, + DtBias: T.Tensor([HV, K], dtype=torch.float32) = None, ): with T.Kernel(T.ceildiv(V, v_per_cta), B * HV, threads=num_threads) as (v_start, bhv_idx): + if fuse_kda_gate: + T.import_source('extern "C" __device__ float __nv_expf(float);\n') tidx = T.get_thread_binding(0) b_id = bhv_idx // HV hv_id = bhv_idx % HV @@ -493,13 +500,30 @@ def fused_recurrent_gated_delta_rule_transposed_main( q_smem = T.alloc_shared([SEQLEN, K], T.float32) k_smem = T.alloc_shared([SEQLEN, K], T.float32) g_exp_smem = T.alloc_shared([SEQLEN], T.float32) + if use_shared_token_inputs or fuse_kda_gate: beta_smem = T.alloc_shared([SEQLEN], T.float32) + if fuse_kda_gate: + kda_decay_smem = T.alloc_shared([SEQLEN, K], T.float32) + if state_id >= 0 and state_id < N: if use_shared_token_inputs: precompute_shared_token_inputs(Query, Key, G, Beta, q_smem, k_smem, g_exp_smem, beta_smem, b_id, h_id, hv_id, warp_id, k_off, K, SEQLEN, k_per_thr, scale, use_qk_l2norm_in_kernel, use_g, use_beta) + elif fuse_kda_gate: + # Preserve PyTorch sigmoid rounding; reuse gates across state waves. + for s in T.Parallel(SEQLEN): + beta_smem[s] = T.ieee_frcp(1.0 + T.call_extern( + 'float32', '__nv_expf', -T.cast(Beta[b_id, s, hv_id], T.float32))) + for s, c in T.Parallel(SEQLEN, K): + gate_value = T.alloc_var(T.float32) + gate_value = T.cast(G[b_id, s, hv_id, c], T.float32) + if has_dt_bias: + gate_value += DtBias[hv_id, c] + gate_value = kda_lower_bound / (1.0 + T.exp(-T.exp(ALog[hv_id]) * gate_value)) + kda_decay_smem[s, c] = T.exp(gate_value) + T.sync_threads() for wave_id in range(num_waves): v_warp_off = wave_id * num_warps * v_per_warp + warp_id * v_per_warp @@ -526,6 +550,9 @@ def fused_recurrent_gated_delta_rule_transposed_main( if use_shared_token_inputs: g_exp = g_exp_smem[seq_id] beta = beta_smem[seq_id] + elif fuse_kda_gate: + g_exp = 1.0 + beta = beta_smem[seq_id] else: if use_g and not channelwise_g: if lane_id == 0: @@ -550,7 +577,10 @@ def fused_recurrent_gated_delta_rule_transposed_main( if channelwise_g: g_local = T.alloc_local([k_per_thr], T.float32) for j in T.Unroll(k_per_thr): - g_local[j] = T.exp(T.cast(G[b_id, seq_id, hv_id, k_off + j], T.float32)) + if fuse_kda_gate: + g_local[j] = kda_decay_smem[seq_id, k_off + j] + else: + g_local[j] = T.exp(T.cast(G[b_id, seq_id, hv_id, k_off + j], T.float32)) update_recurrent_state(h_local, k_local, v_local, g_local, beta, k_per_thr, v_per_warp, channelwise_g=True) else: @@ -597,6 +627,9 @@ def fused_recurrent_gated_delta_rule( state_indices: torch.Tensor | None = None, cache_seqlens: torch.Tensor | None = None, transpose_state_layout: bool = False, + a_log: torch.Tensor | None = None, + dt_bias: torch.Tensor | None = None, + lower_bound: float | None = None, ) -> tuple[torch.Tensor, torch.Tensor | None]: """Fused recurrent gated delta rule. @@ -623,6 +656,10 @@ def fused_recurrent_gated_delta_rule( batch element transpose_state_layout: whether recurrent state is stored as [V, K] instead of [K, V] + a_log: Optional FP32 KDA parameter [HV]. When provided, g and beta + are raw logits and the bounded gate and beta sigmoid are fused. + dt_bias: Optional FP32 KDA gate bias [HV * K] or [HV, K]. + lower_bound: Negative lower bound required for fused KDA gating. Returns: o: [B, T, HV, V] final_state: Recurrent state if ``output_final_state`` is True, @@ -638,6 +675,16 @@ def fused_recurrent_gated_delta_rule( g_dtype = torch.float32 beta_dtype = torch.float32 channelwise_g = g is not None and g.ndim == 4 + fuse_kda_gate = a_log is not None + if fuse_kda_gate: + if not channelwise_g or not transpose_state_layout or beta is None or lower_bound is None or lower_bound >= 0: + raise ValueError('Fused KDA gating requires channelwise raw g/beta, transposed state and a negative bound.') + assert a_log.dtype == torch.float32 and a_log.numel() == HV and a_log.is_contiguous() + if dt_bias is not None: + assert dt_bias.dtype == torch.float32 and dt_bias.numel() == HV * K and dt_bias.is_contiguous() + dt_bias = dt_bias.view(HV, K) + elif dt_bias is not None or lower_bound is not None: + raise ValueError('dt_bias and lower_bound require a_log for fused KDA gating.') if channelwise_g: assert transpose_state_layout, 'Channelwise decay requires transposed state layout' assert g.shape == (*q.shape[:2], HV, K) @@ -712,9 +759,15 @@ def fused_recurrent_gated_delta_rule( transpose_state_layout=transpose_state_layout, num_warps=num_warps, channelwise_g=channelwise_g, + fuse_kda_gate=fuse_kda_gate, + kda_lower_bound=lower_bound if fuse_kda_gate else -5.0, + has_dt_bias=dt_bias is not None, ) - kernel(q, k, v, o, g, beta, final_state, state_indices, cache_seqlens) + if transpose_state_layout: + kernel(q, k, v, o, g, beta, final_state, state_indices, cache_seqlens, a_log, dt_bias) + else: + kernel(q, k, v, o, g, beta, final_state, state_indices, cache_seqlens) if not output_final_state: final_state = None diff --git a/tests/pytorch/kernel/test_kda_fused_gate.py b/tests/pytorch/kernel/test_kda_fused_gate.py new file mode 100644 index 0000000000..2b035402a2 --- /dev/null +++ b/tests/pytorch/kernel/test_kda_fused_gate.py @@ -0,0 +1,120 @@ +# Copyright (c) OpenMMLab. All rights reserved. +import pytest +import torch + + +@pytest.mark.parametrize('steps', [1, 3, 6]) +@pytest.mark.parametrize('state_dtype', [torch.float32, torch.bfloat16]) +@pytest.mark.parametrize('with_bias', [False, True]) +def test_fused_kda_gate_preserves_recurrence_and_dummy_state(steps, state_dtype, with_bias): + from fla.ops.kda.gate import kda_gate_fwd + + from lmdeploy.pytorch.kernels.cuda.gated_delta_rule import fused_recurrent_gated_delta_rule as run + + torch.manual_seed(7) + batch, heads, dim, ring = 3, 2, 128, 6 + mixed = torch.randn(batch, steps, 3 * heads * dim, device='cuda', dtype=torch.bfloat16) + q, k, v = [x.unflatten(-1, (heads, dim)) for x in mixed.chunk(3, dim=-1)] + raw_gate = torch.randn_like(q).contiguous() + raw_beta = torch.randn(batch, steps, heads, device='cuda', dtype=torch.bfloat16) + a_log = torch.randn(heads, device='cuda') + dt_bias = torch.randn(heads * dim, device='cuda') if with_bias else None + initial = torch.randn(batch, ring, heads, dim, dim, device='cuda', dtype=state_dtype) * 0.1 + state = initial.clone() + kwargs = dict(state_indices=torch.tensor([2, -1, 0], device='cuda'), + cache_seqlens=torch.tensor([5, 0, 3], device='cuda', dtype=torch.int32), + output_final_state=True, transpose_state_layout=True, use_qk_l2norm_in_kernel=True) + gate = kda_gate_fwd(raw_gate, a_log, dt_bias, lower_bound=-5.0) + expected, expected_state = run(q, k, v, g=gate, beta=raw_beta.float().sigmoid(), + initial_state=initial.clone(), **kwargs) + actual, _ = run(q, k, v, g=raw_gate, beta=raw_beta, a_log=a_log, dt_bias=dt_bias, + lower_bound=-5.0, initial_state=state, **kwargs) + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + torch.testing.assert_close(state, expected_state, rtol=0, atol=0) + torch.testing.assert_close(state[1], initial[1], rtol=0, atol=0) + assert torch.count_nonzero(actual[1]) == 0 + + +def test_fused_kda_gate_cuda_graph_replay(): + from lmdeploy.pytorch.kernels.cuda.gated_delta_rule import fused_recurrent_gated_delta_rule as run + + torch.manual_seed(29) + q, k, v, raw_gate = [torch.randn(1, 6, 2, 128, device='cuda', dtype=torch.bfloat16) for _ in range(4)] + raw_beta = torch.randn(1, 6, 2, device='cuda', dtype=torch.bfloat16) + initial = torch.randn(1, 6, 2, 128, 128, device='cuda') + state = initial.clone() + history = torch.tensor([5], device='cuda', dtype=torch.int32) + kwargs = dict(g=raw_gate, beta=raw_beta, a_log=torch.randn(2, device='cuda'), + dt_bias=torch.randn(256, device='cuda'), lower_bound=-5.0, cache_seqlens=history, + output_final_state=True, transpose_state_layout=True, use_qk_l2norm_in_kernel=True) + run(q, k, v, initial_state=state, **kwargs) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + actual, _ = run(q, k, v, initial_state=state, **kwargs) + raw_gate.normal_() + raw_beta.normal_() + history.fill_(2) + expected, expected_state = run(q, k, v, initial_state=initial.clone(), **kwargs) + state.copy_(initial) + graph.replay() + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + torch.testing.assert_close(state, expected_state, rtol=0, atol=0) + + +def test_kda_unbounded_gate_keeps_existing_arithmetic(): + from lmdeploy.pytorch.backends.cuda.kda import CudaKdaImpl + + torch.manual_seed(43) + impl = CudaKdaImpl() + q, k, v, raw_gate = [torch.randn(1, 1, 2, 128, device='cuda', dtype=torch.bfloat16) for _ in range(4)] + raw_beta = torch.randn(1, 1, 2, device='cuda', dtype=torch.bfloat16) + a_log = torch.randn(2, device='cuda') + dt_bias = torch.randn(256, device='cuda') + initial = torch.randn(1, 2, 128, 128, device='cuda') + gate = impl.kda_gate(raw_gate, a_log, dt_bias) + expected, expected_state = impl.recurrent_func( + q, k, v, g=gate, beta=raw_beta.float().sigmoid(), initial_state=initial.clone(), + output_final_state=True, transpose_state_layout=True, use_qk_l2norm_in_kernel=True) + actual, state = impl._decode_recurrent( + q, k, v, raw_gate, raw_beta, a_log, dt_bias, initial.clone(), lower_bound=None) + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + torch.testing.assert_close(state, expected_state, rtol=0, atol=0) + + +@pytest.mark.parametrize('beta_dtype', [torch.bfloat16, torch.float32]) +def test_fused_kda_gate_extreme_beta(beta_dtype): + from fla.ops.kda.gate import kda_gate_fwd + + from lmdeploy.pytorch.kernels.cuda.gated_delta_rule import fused_recurrent_gated_delta_rule as run + + torch.manual_seed(53) + q, k, v, raw_gate = [torch.randn(9, 1, 2, 128, device='cuda', dtype=torch.bfloat16) for _ in range(4)] + raw_beta = torch.tensor([-100, -89, -88, -87.5, -86, 0, 86, 88, 100], device='cuda', dtype=beta_dtype) + raw_beta = raw_beta[:, None, None].expand(9, 1, 2).contiguous() + a_log = torch.randn(2, device='cuda') + initial = torch.randn(9, 2, 128, 128, device='cuda') + kwargs = dict(output_final_state=True, transpose_state_layout=True, use_qk_l2norm_in_kernel=True) + expected, expected_state = run( + q, k, v, g=kda_gate_fwd(raw_gate, a_log, lower_bound=-5.0), beta=raw_beta.float().sigmoid(), + initial_state=initial.clone(), **kwargs) + actual, state = run(q, k, v, g=raw_gate, beta=raw_beta, a_log=a_log, lower_bound=-5.0, + initial_state=initial.clone(), **kwargs) + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + torch.testing.assert_close(state, expected_state, rtol=0, atol=0) + + +@pytest.mark.parametrize('invalid', [ + 'missing_bound', 'positive_bound', 'missing_beta', 'untransposed', 'missing_a_log' +]) +def test_fused_kda_gate_rejects_unsupported_contract(invalid): + from lmdeploy.pytorch.kernels.cuda.gated_delta_rule import fused_recurrent_gated_delta_rule as run + + q = torch.empty(1, 1, 2, 128) + kwargs = dict(g=q, beta=torch.empty(1, 1, 2), a_log=torch.zeros(2), + dt_bias=torch.zeros(256), lower_bound=-5.0, transpose_state_layout=True) + overrides = dict(missing_bound=dict(lower_bound=None), positive_bound=dict(lower_bound=1.0), + missing_beta=dict(beta=None), untransposed=dict(transpose_state_layout=False), + missing_a_log=dict(a_log=None)) + kwargs.update(overrides[invalid]) + with pytest.raises(ValueError, match='KDA gating'): + run(q, q, q, initial_state=torch.empty(1, 2, 128, 128), **kwargs) From 9d8c5f1b6fe582b4bc754701608edec2eeb2f9ab Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Wed, 23 Sep 2026 14:14:48 +0000 Subject: [PATCH 15/39] revert: drop DeepGEMM masked GEMM alias compatibility from GLM PR --- .../pytorch/third_party/deep_gemm/__init__.py | 31 +++++++++---------- 1 file changed, 14 insertions(+), 17 deletions(-) diff --git a/lmdeploy/pytorch/third_party/deep_gemm/__init__.py b/lmdeploy/pytorch/third_party/deep_gemm/__init__.py index 336f13634f..cc9cb4333c 100644 --- a/lmdeploy/pytorch/third_party/deep_gemm/__init__.py +++ b/lmdeploy/pytorch/third_party/deep_gemm/__init__.py @@ -78,23 +78,20 @@ def m_grouped_fp8_gemm_nt_contiguous(a, b, d, m_indices, recipe=None, compiled_d try: from deep_gemm import m_grouped_fp8_gemm_nt_masked except Exception: - try: - from deep_gemm import fp8_m_grouped_gemm_nt_masked as m_grouped_fp8_gemm_nt_masked - except Exception: - from deep_gemm import m_grouped_gemm_fp8_fp8_bf16_nt_masked - - def m_grouped_fp8_gemm_nt_masked(a, - b, - d, - masked_m, - expected_m, - recipe=None, - compiled_dims='nk', - disable_ue8m0_cast=False): - assert recipe is None - assert compiled_dims == 'nk' - assert disable_ue8m0_cast is False - return m_grouped_gemm_fp8_fp8_bf16_nt_masked(a, b, d, masked_m, expected_m) + from deep_gemm import m_grouped_gemm_fp8_fp8_bf16_nt_masked + + def m_grouped_fp8_gemm_nt_masked(a, + b, + d, + masked_m, + expected_m, + recipe=None, + compiled_dims='nk', + disable_ue8m0_cast=False): + assert recipe is None + assert compiled_dims == 'nk' + assert disable_ue8m0_cast is False + return m_grouped_gemm_fp8_fp8_bf16_nt_masked(a, b, d, masked_m, expected_m) try: From 7e0b9dedba753d8762504a1b972f6a745c7737de Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Wed, 23 Sep 2026 14:26:57 +0000 Subject: [PATCH 16/39] perf: batch GLM KPool prefill cache updates on device --- lmdeploy/pytorch/backends/cuda/kpool.py | 29 ++- .../pytorch/kernels/cuda/fill_kv_cache.py | 69 +++++++ lmdeploy/pytorch/kernels/cuda/kpool.py | 169 ++++++++++++++++++ lmdeploy/pytorch/models/glm5_next.py | 12 ++ tests/pytorch/kernel/test_kpool_prefill.py | 164 +++++++++++++++++ 5 files changed, 442 insertions(+), 1 deletion(-) create mode 100644 lmdeploy/pytorch/kernels/cuda/kpool.py create mode 100644 tests/pytorch/kernel/test_kpool_prefill.py diff --git a/lmdeploy/pytorch/backends/cuda/kpool.py b/lmdeploy/pytorch/backends/cuda/kpool.py index 400526f32f..b5bc5d1db9 100644 --- a/lmdeploy/pytorch/backends/cuda/kpool.py +++ b/lmdeploy/pytorch/backends/cuda/kpool.py @@ -8,11 +8,38 @@ import torch from torch import Tensor +from lmdeploy.pytorch.kernels.cuda.fill_kv_cache import fill_indexed_key_cache +from lmdeploy.pytorch.kernels.cuda.kpool import compress_kpool, partition_kpool from lmdeploy.pytorch.kernels.cuda.sparse_index_topk import ( is_sparse_index_topk_supported, sparse_index_topk, ) -from lmdeploy.pytorch.nn.kpool import kpool_compress, kpool_quantize_fp8 +from lmdeploy.pytorch.nn.kpool import ( + kpool_compress, + kpool_packed_cache_views, + kpool_quantize_fp8, +) + +from .gated_delta_rule import _state_scatter + + +def kpool_prefill_update_cuda(keys, scores, tail_keys, tail_scores, state_ids, + q_seqlens, kv_seqlens, packed_cache, block_offsets, + ape, pool_size, round_scale): + """Batch ragged pool assembly without a host read or cache-arena copy.""" + closed_keys, closed_scores, group_ids, requests, valid, next_keys, next_scores = partition_kpool( + keys, scores, tail_keys, tail_scores, state_ids, q_seqlens, kv_seqlens, pool_size) + if closed_keys.size(0): + compress = (compress_kpool if torch.cuda.get_device_capability(keys.device)[0] >= 9 + else kpool_compress_quantize_cuda) + values, scales = compress( + closed_keys, closed_scores, ape, mode='extend', round_scale=round_scale) + cache_keys, cache_scales = kpool_packed_cache_views(packed_cache, keys.size(-1)) + fill_indexed_key_cache(values, scales, group_ids, valid, block_offsets, + cache_keys, cache_scales, page_step=pool_size, request_ids=requests) + slots = torch.zeros_like(state_ids) + _state_scatter(tail_keys.unsqueeze(1), state_ids, slots, next_keys) + _state_scatter(tail_scores.unsqueeze(1), state_ids, slots, next_scores) @functools.lru_cache diff --git a/lmdeploy/pytorch/kernels/cuda/fill_kv_cache.py b/lmdeploy/pytorch/kernels/cuda/fill_kv_cache.py index e96819761c..d85f19e0f4 100644 --- a/lmdeploy/pytorch/kernels/cuda/fill_kv_cache.py +++ b/lmdeploy/pytorch/kernels/cuda/fill_kv_cache.py @@ -19,6 +19,75 @@ Q_POLICY_TURBO = tl.constexpr(42) +@triton.jit +def _fill_indexed_key_cache_kernel( + Keys, Scales, GroupIds, Valid, RequestIds, BlockOffsets, KeyCache, ScaleCache, + BATCH: tl.constexpr, COLUMNS: tl.constexpr, PAGE: tl.constexpr, + PAGE_STEP: tl.constexpr, WIDTH: tl.constexpr, + stride_kr: tl.constexpr, stride_kd: tl.constexpr, + stride_sr: tl.constexpr, stride_br: tl.constexpr, stride_bc: tl.constexpr, + stride_kcb: tl.constexpr, stride_kcs: tl.constexpr, stride_kcd: tl.constexpr, + stride_scb: tl.constexpr, stride_scs: tl.constexpr, + BLOCK_D: tl.constexpr, HAS_REQUEST_IDS: tl.constexpr, +): + row = tl.program_id(0) + if tl.load(Valid + row): + group = tl.load(GroupIds + row).to(tl.int64) + column = tl.minimum(tl.maximum(group // PAGE * PAGE_STEP, 0), COLUMNS - 1) + request = tl.load(RequestIds + row) if HAS_REQUEST_IDS else row % BATCH + page = tl.load(BlockOffsets + request * stride_br + column * stride_bc).to(tl.int64) + slot = group % PAGE + d = tl.arange(0, BLOCK_D) + key = tl.load(Keys + row * stride_kr + d * stride_kd, d < WIDTH, other=0.0) + scale = tl.load(Scales + row * stride_sr) + tl.store(KeyCache + page * stride_kcb + slot * stride_kcs + d * stride_kcd, key, d < WIDTH) + tl.store(ScaleCache + page * stride_scb + slot * stride_scs, scale) + + +def fill_indexed_key_cache(keys: Tensor, scales: Tensor, group_ids: Tensor, + valid: Tensor, block_offsets: Tensor, + key_cache: Tensor, scale_cache: Tensor, + page_step: int = 1, request_ids: Tensor | None = None) -> None: + """Scatter prequantized index keys and scales without reading the arena. + + Rows are step-major [steps * batch] unless request_ids is provided. + Each valid destination must be unique; + inactive rows perform no load/store on the cache. ``page_step`` maps a + compressed page to its column in the uncompressed token page table. + Cache views retain their actual strides (including packed DSA storage). + """ + rows, width = keys.shape + batch, columns = block_offsets.shape + if batch == 0 or columns == 0 or (request_ids is None and rows % batch) or page_step < 1: + raise ValueError('Invalid indexed-cache batch/page geometry.') + if group_ids.shape != (rows,) or valid.shape != (rows,): + raise ValueError('One group id and valid flag are required per key.') + if request_ids is not None and request_ids.shape != (rows,): + raise ValueError('One request id is required per key.') + if scales.shape not in ((rows,), (rows, 1)): + raise ValueError('One scale is required per key.') + if keys.dtype != key_cache.dtype or scales.dtype != scale_cache.dtype: + raise TypeError('Prequantized keys/scales must match cache dtypes.') + if key_cache.ndim != 3 or key_cache.size(-1) != width: + raise ValueError('Expected [pages, entries, width] key cache.') + if scale_cache.shape != (*key_cache.shape[:2], 1): + raise ValueError('Expected one scale per cache entry.') + if rows == 0: + return + if keys.dtype == torch.float8_e4m3fn: + # Copy the quantized representation even on devices without native FP8. + keys = keys.view(torch.uint8) + key_cache = key_cache.view(torch.uint8) + _fill_indexed_key_cache_kernel[(rows,)]( + keys, scales, group_ids.contiguous(), valid.contiguous(), + request_ids.contiguous() if request_ids is not None else None, + block_offsets, key_cache, scale_cache, + batch, columns, key_cache.size(1), page_step, width, + *keys.stride(), scales.stride(0), *block_offsets.stride(), + *key_cache.stride(), *scale_cache.stride()[:2], + triton.next_power_of_2(width), request_ids is not None, num_warps=4) + + @triton.jit def _quant_int8(val): val_min = tl.min(val, 1) diff --git a/lmdeploy/pytorch/kernels/cuda/kpool.py b/lmdeploy/pytorch/kernels/cuda/kpool.py new file mode 100644 index 0000000000..cb0cd4be37 --- /dev/null +++ b/lmdeploy/pytorch/kernels/cuda/kpool.py @@ -0,0 +1,169 @@ +# Copyright (c) OpenMMLab. All rights reserved. +import torch +import triton +import triton.language as tl +import triton.language.extra.cuda.libdevice as libdevice + + +@triton.jit +def _partition_kpool_kernel( + Keys, Scores, TailKeys, TailScores, StateIds, QLens, KVLens, + ClosedKeys, ClosedScores, GroupIds, RequestIds, Valid, NextKeys, NextScores, + BATCH: tl.constexpr, CAPACITY: tl.constexpr, STATES: tl.constexpr, + POOL: tl.constexpr, WIDTH: tl.constexpr, BLOCK_B: tl.constexpr, BLOCK_D: tl.constexpr, + STRIDE_TK: tl.constexpr, STRIDE_TS: tl.constexpr, +): + row = tl.program_id(0) + slots = tl.arange(0, POOL) + d = tl.arange(0, BLOCK_D) + group = tl.full((), 0, tl.int64) + if row < CAPACITY: + batches = tl.arange(0, BLOCK_B) + q = tl.load(QLens + batches, batches < BATCH, other=0) + kv = tl.load(KVLens + batches, batches < BATCH, other=0) + counts = ((kv - q) % POOL + q) // POOL + ends = tl.cumsum(counts) + request = tl.minimum(tl.sum(((row >= ends) & (batches < BATCH)).to(tl.int32)), BATCH - 1) + group_start = tl.sum(tl.where(batches == request, ends - counts, 0)) + token_start = tl.sum(tl.where(batches < request, q, 0)) + valid = row < tl.sum(counts) + q_len = tl.load(QLens + request) + kv_len = tl.load(KVLens + request) + history = kv_len - q_len + offsets = (row - group_start) * POOL + slots - history % POOL + group = (history // POOL + row - group_start).to(tl.int64) + else: + request = row - CAPACITY + batches = tl.arange(0, BLOCK_B) + q = tl.load(QLens + batches, batches < request, other=0) + token_start = tl.sum(q) + q_len = tl.load(QLens + request) + kv_len = tl.load(KVLens + request) + history = kv_len - q_len + offsets = ((history % POOL + q_len) // POOL) * POOL + slots - history % POOL + valid = True + state_id = tl.load(StateIds + request) + valid = valid & (state_id >= 0) & (state_id < STATES) + active_slots = tl.full((POOL,), True, tl.int1) + if row >= CAPACITY: + active_slots = slots < kv_len % POOL + prior_mask = valid & active_slots & (offsets < 0) + token_mask = valid & active_slots & (offsets >= 0) & (offsets < q_len) + prior_offsets = offsets + history % POOL + old_k = tl.load(TailKeys + state_id * STRIDE_TK + prior_offsets[:, None] * WIDTH + d[None, :], + prior_mask[:, None] & (d[None, :] < WIDTH), other=0) + old_s = tl.load(TailScores + state_id * STRIDE_TS + prior_offsets[:, None] * WIDTH + d[None, :], + prior_mask[:, None] & (d[None, :] < WIDTH), other=0) + new_k = tl.load(Keys + (token_start + offsets[:, None]) * WIDTH + d[None, :], + token_mask[:, None] & (d[None, :] < WIDTH), other=0) + new_s = tl.load(Scores + (token_start + offsets[:, None]) * WIDTH + d[None, :], + token_mask[:, None] & (d[None, :] < WIDTH), other=0) + key = tl.where((offsets < 0)[:, None], old_k, new_k) + score = tl.where((offsets < 0)[:, None], old_s, new_s) + if row < CAPACITY: + tl.store(ClosedKeys + (row * POOL + slots[:, None]) * WIDTH + d[None, :], key, d[None, :] < WIDTH) + tl.store(ClosedScores + (row * POOL + slots[:, None]) * WIDTH + d[None, :], score, d[None, :] < WIDTH) + tl.store(GroupIds + row, group) + tl.store(RequestIds + row, request) + tl.store(Valid + row, valid) + else: + tl.store(NextKeys + (request * POOL + slots[:, None]) * WIDTH + d[None, :], key, d[None, :] < WIDTH) + tl.store(NextScores + (request * POOL + slots[:, None]) * WIDTH + d[None, :], score, d[None, :] < WIDTH) + + +def partition_kpool(keys, scores, tail_keys, tail_scores, state_ids, q_seqlens, kv_seqlens, pool_size): + """Assemble ragged closed pools and next tails without host metadata reads. + + Group capacity depends only on input shapes. Invalid group rows are zero padded and masked; persistent state is read + but never modified here. + """ + batch = q_seqlens.numel() + tokens, width = keys.shape + if pool_size <= 1 or pool_size & (pool_size - 1) or not batch or scores.shape != keys.shape: + raise ValueError('Expected matching keys/scores, a nonempty batch and a power-of-two pool size.') + if state_ids.shape != (batch,) or kv_seqlens.shape != (batch,): + raise ValueError('Expected one state id and KV length per request.') + if tail_keys.shape != tail_scores.shape or tail_keys.shape[1:] != (pool_size, width): + raise ValueError('Expected matching [states, pool_size, width] tail caches.') + if tail_keys.stride()[1:] != (width, 1) or tail_scores.stride()[1:] != (width, 1): + raise ValueError('Tail cache rows must be contiguous.') + capacity = (tokens + batch * (pool_size - 1)) // pool_size + closed_keys = keys.new_empty((capacity, pool_size, width), dtype=torch.promote_types(keys.dtype, tail_keys.dtype)) + closed_scores = scores.new_empty((capacity, pool_size, width), + dtype=torch.promote_types(scores.dtype, tail_scores.dtype)) + groups = state_ids.new_empty(capacity) + requests = state_ids.new_empty(capacity) + valid = torch.empty(capacity, device=keys.device, dtype=torch.bool) + next_keys = tail_keys.new_empty((batch, pool_size, width)) + next_scores = tail_scores.new_empty((batch, pool_size, width)) + _partition_kpool_kernel[(capacity + batch,)]( + keys.contiguous(), scores.contiguous(), tail_keys, tail_scores, + state_ids.contiguous(), q_seqlens.contiguous(), kv_seqlens.contiguous(), + closed_keys, closed_scores, groups, requests, valid, next_keys, next_scores, + batch, capacity, tail_keys.size(0), pool_size, width, triton.next_power_of_2(batch), + triton.next_power_of_2(width), tail_keys.stride(0), tail_scores.stride(0), num_warps=4) + return closed_keys, closed_scores, groups, requests, valid, next_keys, next_scores + + +@triton.jit +def _compress_kpool_kernel(K, S, A, O, Scale, WIDTH: tl.constexpr, POOL: tl.constexpr, + ONLINE: tl.constexpr, ROUND: tl.constexpr, LEVELS: tl.constexpr): + row = tl.program_id(0) + d = tl.arange(0, WIDTH) + maximum = tl.full((WIDTH,), -float('inf'), tl.float32) + denominator = tl.full((WIDTH,), 0, tl.float32) + accumulator = tl.full((WIDTH,), 0, tl.float32) + if not ONLINE: + for slot in tl.static_range(POOL): + score = tl.load(S + (row * POOL + slot) * WIDTH + d).to(tl.float32) + score += tl.load(A + slot * WIDTH + d).to(tl.float32) + maximum = tl.maximum(maximum, score) + for slot in tl.static_range(POOL): + score = tl.load(S + (row * POOL + slot) * WIDTH + d).to(tl.float32) + score += tl.load(A + slot * WIDTH + d).to(tl.float32) + if ONLINE: + new_maximum = tl.maximum(maximum, score) + rescale = libdevice.exp(maximum - new_maximum) + denominator = denominator * rescale + accumulator = accumulator * rescale + maximum = new_maximum + probability = libdevice.exp(score - maximum) + key = tl.load(K + (row * POOL + slot) * WIDTH + d).to(tl.float32) + denominator = denominator + probability + accumulator = accumulator + key * probability + value = tl.div_rn(accumulator, denominator).to(tl.bfloat16).to(tl.float32) + for level in tl.static_range(LEVELS): + stride = 1 << level + other = tl.gather(value, d ^ stride, axis=0) + value = tl.where((d & stride) == 0, value + other, other - value) + value = (value * (WIDTH**-0.5)).to(tl.bfloat16).to(tl.float32) + scale = tl.maximum(tl.max(tl.abs(value), 0), 1e-4) * (1.0 / 448.0) + if ROUND: + scale = libdevice.exp2(libdevice.ceil(libdevice.log2(scale))) + quant = tl.minimum(tl.maximum(tl.div_rn(value, scale), -448.0), 448.0) + tl.store(O + row * WIDTH + d, quant.to(tl.float8e4nv)) + tl.store(Scale + row, scale) + + +def compress_kpool(keys: torch.Tensor, scores: torch.Tensor, ape: torch.Tensor, + *, mode: str, round_scale: bool) -> tuple[torch.Tensor, torch.Tensor]: + """Fuse weighted pooling, BF16 Hadamard rotation and one-block FP8 + quantization. + + The online/two-pass reduction order and both BF16 round trips match the reference KPool operation. Precise libdevice + functions and disabled FMA fusion preserve its FP32 arithmetic boundaries. + """ + if mode not in ('extend', 'decode'): + raise ValueError(f'Unsupported pool compression mode: {mode}') + if keys.ndim != 3 or scores.shape != keys.shape or ape.shape != keys.shape[1:]: + raise ValueError('Expected matching [groups, pool, width] keys/scores and [pool, width] APE.') + groups, pool, width = keys.shape + if width <= 0 or width & (width - 1): + raise ValueError('Pool width must be a positive power of two.') + out = torch.empty(groups, width, device=keys.device, dtype=torch.float8_e4m3fn) + scale = torch.empty(groups, 1, device=keys.device, dtype=torch.float32) + if groups: + _compress_kpool_kernel[(groups,)]( + keys.contiguous(), scores.contiguous(), ape.contiguous(), out, scale, width, pool, + mode == 'extend', round_scale, width.bit_length() - 1, num_warps=4, enable_fp_fusion=False) + return out, scale diff --git a/lmdeploy/pytorch/models/glm5_next.py b/lmdeploy/pytorch/models/glm5_next.py index 9ab81f1ac2..deeba4a7af 100644 --- a/lmdeploy/pytorch/models/glm5_next.py +++ b/lmdeploy/pytorch/models/glm5_next.py @@ -18,6 +18,7 @@ ) from lmdeploy.pytorch.backends.cuda.kpool import ( kpool_compress_quantize_cuda, + kpool_prefill_update_cuda, kpool_score_contiguous_cuda, kpool_score_paged_cuda, kpool_select_groups_cuda, @@ -919,6 +920,17 @@ def save_ring(lengths): save_ring(history_lengths + step + 1) return indexer_k_cache + if key.is_cuda: + prefill_ids = state_ids if ring_states is None else torch.where(valid_requests, state_ids, -1) + kpool_prefill_update_cuda( + key, score, tail_k_state, tail_score_state, prefill_ids, + attn_metadata.q_seqlens, attn_metadata.kv_seqlens, + indexer_k_cache, attn_metadata.block_offsets, + self.indexer.index_kpool_compress_ape, self.index_kpool, + self.indexer.scale_fmt is not None) + save_ring(attn_metadata.kv_seqlens) + return indexer_k_cache + q_seqlens = attn_metadata.q_seqlens.tolist() kv_seqlens = attn_metadata.kv_seqlens.tolist() if len(q_seqlens) != len(kv_seqlens): diff --git a/tests/pytorch/kernel/test_kpool_prefill.py b/tests/pytorch/kernel/test_kpool_prefill.py new file mode 100644 index 0000000000..c060ce787f --- /dev/null +++ b/tests/pytorch/kernel/test_kpool_prefill.py @@ -0,0 +1,164 @@ +# Copyright (c) OpenMMLab. All rights reserved. +import pytest +import torch + + +def reference_update(keys, scores, states, ids, q_lens, kv_lens, cache, blocks, ape, round_scale): + from lmdeploy.pytorch.backends.cuda.kpool import kpool_compress_quantize_cuda + from lmdeploy.pytorch.nn.kpool import kpool_partition_update, kpool_write_packed_cache + + start = 0 + for request, (q_len, kv_len, state_id) in enumerate(zip(q_lens.tolist(), kv_lens.tolist(), ids.tolist())): + if state_id >= 0: + history = kv_len - q_len + old_tail = history % 4 + update = kpool_partition_update(keys[start:start + q_len], scores[start:start + q_len], history, 4, + states[0][state_id, :old_tail], states[1][state_id, :old_tail]) + if update.closed_group_ids.numel(): + values, scales = kpool_compress_quantize_cuda( + update.closed_keys, update.closed_scores, ape, mode='extend', round_scale=round_scale) + kpool_write_packed_cache(cache, blocks[request], update.closed_group_ids, values, scales, 4) + for state, tail in zip(states, (update.tail_keys, update.tail_scores)): + # Clone because the empty-query reference can alias the source tail. + tail = tail.clone() + state[state_id].zero_() + state[state_id, :tail.size(0)].copy_(tail) + start += q_len + + +def make_case(lengths, histories, round_scale): + torch.manual_seed(17) + batch = len(lengths) + keys = torch.randn(sum(lengths), 128, device='cuda', dtype=torch.bfloat16) + scores = torch.randn_like(keys) + q_lens = torch.tensor(lengths, device='cuda', dtype=torch.int32) + kv_lens = q_lens + torch.tensor(histories, device='cuda', dtype=torch.int32) + ids = torch.arange(batch, 0, -1, device='cuda') + states = tuple(torch.randn(batch + 2, 4, 128, device='cuda', dtype=torch.bfloat16) for _ in range(2)) + blocks = torch.arange(1, batch * 160 + 1, device='cuda').reshape(batch, 160) + cache = torch.randint(0, 256, (batch * 160 + 1, 64, 1, 132), device='cuda', dtype=torch.uint8) + ape = torch.randn(4, 128, device='cuda') + return keys, scores, states, ids, q_lens, kv_lens, cache, blocks, ape, round_scale + + +def candidate_update(keys, scores, states, ids, q_lens, kv_lens, cache, blocks, ape, round_scale): + from lmdeploy.pytorch.backends.cuda.kpool import kpool_prefill_update_cuda + + kpool_prefill_update_cuda(keys, scores, *states, ids, q_lens, kv_lens, cache, blocks, ape, 4, round_scale) + + +@pytest.mark.parametrize('lengths,histories', [ + ([0], [3]), ([0, 0], [0, 3]), ([1, 2], [0, 0]), + ([0, 1, 3, 4, 5], [0, 3, 1, 4, 7]), ([511, 513], [3, 256]), ([8192], [0]), +]) +@pytest.mark.parametrize('round_scale', [False, True]) +@pytest.mark.parametrize('metadata_dtype', [torch.int32, torch.int64]) +def test_kpool_ragged_prefill_matches_reference(lengths, histories, round_scale, metadata_dtype): + args = make_case(lengths, histories, round_scale) + keys, scores, states, ids, q_lens, kv_lens, cache, blocks, ape, _ = args + q_lens, kv_lens = q_lens.to(metadata_dtype), kv_lens.to(metadata_dtype) + expected_states = tuple(state.clone() for state in states) + expected_cache = cache.clone() + reference_update(keys, scores, expected_states, ids, q_lens, kv_lens, expected_cache, blocks, ape, round_scale) + candidate_update(keys, scores, states, ids, q_lens, kv_lens, cache, blocks, ape, round_scale) + torch.testing.assert_close(cache, expected_cache, rtol=0, atol=0) + for actual, expected in zip(states, expected_states): + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + + +def test_kpool_ragged_prefill_graph_changes_layout_and_padding(): + args = make_case([3, 5, 0], [0, 1, 2], True) + keys, scores, states, ids, q_lens, kv_lens, cache, blocks, ape, _ = args + ids[-1] = -1 + initial_states = tuple(state.clone() for state in states) + initial_cache = cache.clone() + candidate_update(*args) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + candidate_update(*args) + for lengths, histories, reordered in [([1, 0, 7], [5, 4, 8], [1, 3, 2]), + ([4, 3, 1], [7, 1, 9], [-1, 2, 1])]: + q_lens.copy_(torch.tensor(lengths, device='cuda')) + kv_lens.copy_(q_lens + torch.tensor(histories, device='cuda')) + ids.copy_(torch.tensor(reordered, device='cuda')) + keys.normal_() + scores.normal_() + expected_states = tuple(state.clone() for state in initial_states) + expected_cache = initial_cache.clone() + reference_update(keys, scores, expected_states, ids, q_lens, kv_lens, expected_cache, blocks, ape, True) + for state, initial in zip(states, initial_states): + state.copy_(initial) + cache.copy_(initial_cache) + graph.replay() + torch.testing.assert_close(cache, expected_cache, rtol=0, atol=0) + for actual, expected in zip(states, expected_states): + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + + +def test_kpool_prefill_layer_views_and_chunk_continuation(): + keys, scores, states, ids, q_lens, kv_lens, cache, blocks, ape, _ = make_case([3, 5], [6, 7], True) + # Runtime states are layer views of [request, layer, pool, width] storage. + banks = tuple(torch.randn(state.size(0), 11, 4, 128, device='cuda', dtype=state.dtype) for state in states) + expected_banks = tuple(bank.clone() for bank in banks) + states = tuple(bank[:, 3] for bank in banks) + expected_states = tuple(bank[:, 3] for bank in expected_banks) + expected_cache = cache.clone() + scores = scores.float() + for lengths in ([3, 5], [1, 7], [6, 2]): + history = kv_lens - q_lens + q_lens.copy_(torch.tensor(lengths, device='cuda')) + kv_lens.copy_(history + q_lens) + keys.normal_() + scores.normal_() + reference_update(keys, scores, expected_states, ids, q_lens, kv_lens, expected_cache, blocks, ape, True) + candidate_update(keys, scores, states, ids, q_lens, kv_lens, cache, blocks, ape, True) + torch.testing.assert_close(cache, expected_cache, rtol=0, atol=0) + for actual, expected in zip(banks, expected_banks): + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + kv_lens.add_(q_lens) + + +def test_indexed_key_scatter_step_major_strides_and_mask(): + from lmdeploy.pytorch.kernels.cuda.fill_kv_cache import fill_indexed_key_cache + from lmdeploy.pytorch.nn.kpool import kpool_packed_cache_views + + torch.manual_seed(59) + cache = torch.randint(0, 256, (17, 64, 1, 132), device='cuda', dtype=torch.uint8) + expected = cache.clone() + keys, scales = kpool_packed_cache_views(cache, 128) + ref_keys, ref_scales = kpool_packed_cache_views(expected, 128) + values = torch.randn(128, 6, device='cuda').to(torch.float8_e4m3fn).T + value_scales = torch.rand(12, device='cuda')[::2] + blocks = torch.arange(1, 17, device='cuda').reshape(2, 8) + groups = torch.tensor([0, -1, 1, 2, -1, 3], device='cuda') + valid = groups >= 0 + for row in [0, 2, 3, 5]: + ref_keys[blocks[row % 2, 0], groups[row]] = values[row] + ref_scales[blocks[row % 2, 0], groups[row], 0] = value_scales[row] + fill_indexed_key_cache(values, value_scales, groups, valid, blocks, keys, scales, page_step=4) + torch.testing.assert_close(cache, expected, rtol=0, atol=0) + + +@pytest.mark.skipif(not torch.cuda.is_available() or torch.cuda.get_device_capability()[0] < 9, + reason='Native FP8 compression requires Hopper or newer.') +@pytest.mark.parametrize('width,pool', [(64, 2), (128, 4)]) +@pytest.mark.parametrize('mode', ['extend', 'decode']) +@pytest.mark.parametrize('round_scale', [False, True]) +def test_kpool_compression_preserves_fp8_bytes_and_scales(width, pool, mode, round_scale): + from lmdeploy.pytorch.kernels.cuda.kpool import compress_kpool + from lmdeploy.pytorch.nn.kpool import kpool_compress, kpool_quantize_fp8 + + torch.manual_seed(73) + keys = torch.randn(257, pool, width, device='cuda', dtype=torch.bfloat16) + scores = torch.randn_like(keys, dtype=torch.float32 if mode == 'decode' else torch.bfloat16) + ape = torch.randn(pool, width, device='cuda') + keys[0].zero_() + keys[1].mul_(1e-6) + keys[2].mul_(1000) + scores[3, 0].fill_(90) + scores[3, 1:].fill_(-90) + expected = kpool_quantize_fp8(kpool_compress(keys, scores, ape, mode=mode), + block_size=width, round_scale=round_scale) + actual = compress_kpool(keys, scores, ape, mode=mode, round_scale=round_scale) + torch.testing.assert_close(actual[0].view(torch.uint8), expected[0].view(torch.uint8), rtol=0, atol=0) + torch.testing.assert_close(actual[1], expected[1], rtol=0, atol=0) From e542908a0e089989143d098f848500ca567b7d91 Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Wed, 23 Sep 2026 14:42:29 +0000 Subject: [PATCH 17/39] perf: batch GLM KPool prefill scoring and selection --- lmdeploy/pytorch/backends/cuda/kpool.py | 45 +++++++++++- lmdeploy/pytorch/kernels/cuda/kpool.py | 55 ++++++++++++++ lmdeploy/pytorch/models/glm5_next.py | 9 ++- tests/pytorch/kernel/test_kpool_selection.py | 77 ++++++++++++++++++++ 4 files changed, 183 insertions(+), 3 deletions(-) create mode 100644 tests/pytorch/kernel/test_kpool_selection.py diff --git a/lmdeploy/pytorch/backends/cuda/kpool.py b/lmdeploy/pytorch/backends/cuda/kpool.py index b5bc5d1db9..6e48ec2c9f 100644 --- a/lmdeploy/pytorch/backends/cuda/kpool.py +++ b/lmdeploy/pytorch/backends/cuda/kpool.py @@ -9,19 +9,25 @@ from torch import Tensor from lmdeploy.pytorch.kernels.cuda.fill_kv_cache import fill_indexed_key_cache -from lmdeploy.pytorch.kernels.cuda.kpool import compress_kpool, partition_kpool +from lmdeploy.pytorch.kernels.cuda.flatten_kv_cache import flatten_kv_cache +from lmdeploy.pytorch.kernels.cuda.kpool import compress_kpool, kpool_prefill_metadata, partition_kpool from lmdeploy.pytorch.kernels.cuda.sparse_index_topk import ( is_sparse_index_topk_supported, sparse_index_topk, ) from lmdeploy.pytorch.nn.kpool import ( kpool_compress, + kpool_expand_selected_groups, kpool_packed_cache_views, kpool_quantize_fp8, ) from .gated_delta_rule import _state_scatter +# Fuse the existing integer index expansion instead of materializing its +# [tokens, topk] masks and int64 temporaries separately during batched prefill. +_expand_prefill_groups = torch.compile(kpool_expand_selected_groups, dynamic=True, fullgraph=True) + def kpool_prefill_update_cuda(keys, scores, tail_keys, tail_scores, state_ids, q_seqlens, kv_seqlens, packed_cache, block_offsets, @@ -80,6 +86,41 @@ def kpool_compress_quantize_cuda( ) +def kpool_select_prefill_cuda(query_fp8, query_weight, packed_cache, + q_seqlens, kv_seqlens, block_offsets, kv_flatten_size, + pool_size, topk): + """Score and select all ragged prefill requests without host length + reads.""" + _validate_query(query_fp8, query_weight) + rows = query_fp8.size(0) + if not rows: + return torch.empty((0, topk + pool_size - 1), device=query_fp8.device, dtype=torch.int32) + counts, starts, seq, lengths, query_starts, query_ends = kpool_prefill_metadata( + q_seqlens, kv_seqlens, rows, pool_size) + keys, scales = kpool_packed_cache_views(packed_cache, query_fp8.size(-1)) + blocks = block_offsets[:, ::pool_size].contiguous() + # Sum(floor(kv / pool)) can be up to batch - 1 below floor(sum(kv) / pool). + # Give the shared flatten kernel enough pages to zero this padded tail. + tail_pages = (counts.numel() + keys.size(1) - 1) // keys.size(1) + if tail_pages > blocks.size(1): + blocks = torch.nn.functional.pad(blocks, (0, tail_pages - blocks.size(1))) + # Reuse the same-dtype flatten operator by copying both fields as bytes. + flat_keys, flat_scales = flatten_kv_cache( + keys.view(torch.uint8).unsqueeze(2), scales.view(torch.uint8).unsqueeze(2), + counts, blocks, start_loc=starts, out_size=max(1, kv_flatten_size // pool_size)) + max_groups = block_offsets.size(1) * keys.size(1) // pool_size + logits = _get_deep_gemm().fp8_mqa_logits( + query_fp8.contiguous(), + (flat_keys[0].view(torch.float8_e4m3fn), flat_scales[0].view(torch.float32).flatten()), + query_weight.contiguous(), query_starts, query_ends, + clean_logits=False, max_seqlen_k=max_groups) + # Compressed logits are request-local. The selector masks the unwritten + # suffix using each query's causal length, including zero-length rows. + selected = kpool_select_groups_cuda( + logits, lengths, group_topk=topk // pool_size, max_group_length=max_groups) + return _expand_prefill_groups(selected, lengths, pool_size, topk, seq_lens=seq) + + def kpool_select_groups_cuda( logits: Tensor, group_lengths: Tensor, @@ -132,7 +173,7 @@ def kpool_select_groups_cuda( q_seqlens = torch.ones( logits.size(0), dtype=torch.int32, device=logits.device) return sparse_index_topk( - score_window.contiguous(), + score_window, q_seqlens, lengths.clamp(max=max_group_length), group_topk, diff --git a/lmdeploy/pytorch/kernels/cuda/kpool.py b/lmdeploy/pytorch/kernels/cuda/kpool.py index cb0cd4be37..1238ac41f3 100644 --- a/lmdeploy/pytorch/kernels/cuda/kpool.py +++ b/lmdeploy/pytorch/kernels/cuda/kpool.py @@ -5,6 +5,61 @@ import triton.language.extra.cuda.libdevice as libdevice +@triton.jit +def _prefill_offsets_kernel(Q, KV, Counts, Starts, QEnds, + BATCH: tl.constexpr, POOL: tl.constexpr, BLOCK: tl.constexpr): + request = tl.arange(0, BLOCK) + q = tl.load(Q + request, request < BATCH, other=0).to(tl.int32) + groups = tl.load(KV + request, request < BATCH, other=0).to(tl.int32) // POOL + tl.store(Counts + request, groups, request < BATCH) + tl.store(Starts + request, tl.cumsum(groups) - groups, request < BATCH) + tl.store(QEnds + request, tl.cumsum(q), request < BATCH) + + +@triton.jit +def _prefill_query_metadata_kernel(KV, QEnds, Starts, Seq, Lengths, QueryStarts, QueryEnds, + ROWS: tl.constexpr, BATCH: tl.constexpr, POOL: tl.constexpr, + BLOCK: tl.constexpr): + row = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK) + low = tl.full((BLOCK,), 0, tl.int32) + high = tl.full((BLOCK,), BATCH, tl.int32) + # Upper-bound search skips empty requests without a host ragged loop. + while tl.sum((low < high).to(tl.int32), 0) > 0: + mid = (low + high) // 2 + end = tl.load(QEnds + mid, mid < BATCH, other=2147483647) + active = low < high + right = row >= end + low = tl.where(active & right, mid + 1, low) + high = tl.where(active & ~right, mid, high) + request = tl.minimum(low, BATCH - 1) + seq = tl.load(KV + request).to(tl.int32) + row - tl.load(QEnds + request) + 1 + start = tl.load(Starts + request) + tl.store(Seq + row, seq, row < ROWS) + tl.store(Lengths + row, seq // POOL, row < ROWS) + tl.store(QueryStarts + row, start, row < ROWS) + tl.store(QueryEnds + row, start + seq // POOL, row < ROWS) + + +def kpool_prefill_metadata(q_seqlens, kv_seqlens, rows, pool_size): + """Build compressed-cache offsets and causal query bounds on device.""" + batch = q_seqlens.numel() + counts = torch.empty(batch, device=q_seqlens.device, dtype=torch.int32) + starts = torch.empty_like(counts) + q_ends = torch.empty_like(counts) + seq = torch.empty(rows, device=q_seqlens.device, dtype=torch.int32) + lengths = torch.empty_like(seq) + query_starts = torch.empty_like(seq) + query_ends = torch.empty_like(seq) + _prefill_offsets_kernel[(1,)]( + q_seqlens.contiguous(), kv_seqlens.contiguous(), counts, starts, q_ends, + batch, pool_size, triton.next_power_of_2(batch)) + if rows: + _prefill_query_metadata_kernel[(triton.cdiv(rows, 128),)]( + kv_seqlens.contiguous(), q_ends, starts, seq, lengths, query_starts, query_ends, + rows, batch, pool_size, 128) + return counts, starts, seq, lengths, query_starts, query_ends + + @triton.jit def _partition_kpool_kernel( Keys, Scores, TailKeys, TailScores, StateIds, QLens, KVLens, diff --git a/lmdeploy/pytorch/models/glm5_next.py b/lmdeploy/pytorch/models/glm5_next.py index deeba4a7af..88098cbebe 100644 --- a/lmdeploy/pytorch/models/glm5_next.py +++ b/lmdeploy/pytorch/models/glm5_next.py @@ -22,6 +22,7 @@ kpool_score_contiguous_cuda, kpool_score_paged_cuda, kpool_select_groups_cuda, + kpool_select_prefill_cuda, ) from lmdeploy.pytorch.configurations.glm5_next import is_glm5_kda_layer from lmdeploy.pytorch.consts import ( @@ -1089,7 +1090,13 @@ def _select_kpool_indices_prefill( indexer_k_cache: torch.Tensor, attn_metadata: Any, ) -> torch.Tensor: - """Retain the ragged eager implementation for chunked prefill.""" + """Select request-local pooled history for chunked prefill.""" + if query_fp8.is_cuda: + return kpool_select_prefill_cuda( + query_fp8, query_weight, indexer_k_cache, + attn_metadata.q_seqlens, attn_metadata.kv_seqlens, + attn_metadata.block_offsets, attn_metadata.kv_flatten_size, + self.index_kpool, self.index_topk) q_seqlens = attn_metadata.q_seqlens.tolist() kv_seqlens = attn_metadata.kv_seqlens.tolist() logical_parts = [] diff --git a/tests/pytorch/kernel/test_kpool_selection.py b/tests/pytorch/kernel/test_kpool_selection.py new file mode 100644 index 0000000000..4271b25b12 --- /dev/null +++ b/tests/pytorch/kernel/test_kpool_selection.py @@ -0,0 +1,77 @@ +# Copyright (c) OpenMMLab. All rights reserved. +import pytest +import torch + +pytestmark = pytest.mark.skipif(not torch.cuda.is_available() or torch.cuda.get_device_capability()[0] < 9, + reason='DeepGEMM FP8 scoring requires Hopper or newer.') + + +def reference_selection(query, weight, cache, q_lens, kv_lens, blocks): + from lmdeploy.pytorch.backends.cuda.kpool import kpool_score_contiguous_cuda, kpool_select_groups_cuda + from lmdeploy.pytorch.nn.kpool import kpool_expand_selected_groups, kpool_read_packed_cache + + parts = [] + offset = 0 + for request, (q, kv) in enumerate(zip(q_lens.tolist(), kv_lens.tolist())): + seq = kv - q + torch.arange(1, q + 1, device=query.device) + lengths = seq // 4 + keys, scales = kpool_read_packed_cache(cache, blocks[request], kv // 4, 4) + scores = kpool_score_contiguous_cuda(query[offset:offset + q], weight[offset:offset + q], keys, scales, lengths) + selected = kpool_select_groups_cuda(scores, lengths, group_topk=512, max_group_length=kv // 4) + parts.append(kpool_expand_selected_groups(selected, lengths, 4, 2048, seq_lens=seq)) + offset += q + return torch.cat(parts) + + +def make_case(lengths, histories, tied=False): + from lmdeploy.pytorch.nn.kpool import kpool_packed_cache_views + + torch.manual_seed(97) + batch, rows = len(lengths), sum(lengths) + kv = [q + h for q, h in zip(lengths, histories)] + columns = max(1, (max(kv) + 63) // 64) + blocks = (torch.randperm(batch * columns, device='cuda') + 1).reshape(batch, columns) + cache = torch.empty(batch * columns + 1, 64, 1, 132, device='cuda', dtype=torch.uint8) + keys, scales = kpool_packed_cache_views(cache, 128) + keys.copy_(torch.randn(keys.shape, device='cuda').to(torch.float8_e4m3fn)) + scales.uniform_(0.001, 0.05) + query = torch.randn(rows, 32, 128, device='cuda').to(torch.float8_e4m3fn) + weight = torch.zeros(rows, 32, device='cuda') if tied else torch.rand(rows, 32, device='cuda') + return (query, weight, cache, torch.tensor(lengths, device='cuda'), + torch.tensor(kv, device='cuda'), blocks, sum(kv)) + + +@pytest.mark.parametrize('lengths,histories', [ + ([0], [0]), ([1, 2], [0, 0]), ([0, 13, 17], [0, 2079, 6287]), ([511, 513], [3, 4096]), + ([1, 4, 13], [2048, 8189, 32756]), + ([8192], [0]), ([1] * 256, [2] * 256), ([1] * 256, [3] * 256), +]) +@pytest.mark.parametrize('tied', [False, True]) +def test_kpool_prefill_selection_matches_request_loop(lengths, histories, tied): + pytest.importorskip('deep_gemm') + from lmdeploy.pytorch.backends.cuda.kpool import kpool_select_prefill_cuda + + args = make_case(lengths, histories, tied) + expected = reference_selection(*args[:-1]) + actual = kpool_select_prefill_cuda(*args, 4, 2048) + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + + +def test_kpool_prefill_selection_graph_replays_ragged_metadata(): + pytest.importorskip('deep_gemm') + from lmdeploy.pytorch.backends.cuda.kpool import kpool_select_prefill_cuda + + args = make_case([13, 17, 0], [2079, 6278, 0]) + query, weight, cache, q_lens, kv_lens, blocks, capacity = args + kpool_select_prefill_cuda(*args, 4, 2048) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + actual = kpool_select_prefill_cuda(*args, 4, 2048) + for q, kv in [([0, 11, 19], [0, 4097, 4000]), ([17, 0, 13], [3107, 0, 4013])]: + assert sum(kv) <= capacity + q_lens.copy_(torch.tensor(q, device='cuda')) + kv_lens.copy_(torch.tensor(kv, device='cuda')) + weight.uniform_() + expected = reference_selection(query, weight, cache, q_lens, kv_lens, blocks) + graph.replay() + torch.testing.assert_close(actual, expected, rtol=0, atol=0) From a2800a0e07f70372cfdfc91379e8afb2c6a953f8 Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Wed, 23 Sep 2026 14:59:05 +0000 Subject: [PATCH 18/39] perf: store GLM NoPE latent cache without padding --- lmdeploy/pytorch/configurations/glm5_next.py | 2 +- .../pytorch/kernels/cuda/flatten_kv_cache.py | 2 +- lmdeploy/pytorch/models/glm5_next.py | 8 +- tests/pytorch/config/test_glm5_runtime_env.py | 27 +++++ tests/pytorch/kernel/test_glm5_nope_cache.py | 103 ++++++++++++++++++ 5 files changed, 133 insertions(+), 9 deletions(-) create mode 100644 tests/pytorch/kernel/test_glm5_nope_cache.py diff --git a/lmdeploy/pytorch/configurations/glm5_next.py b/lmdeploy/pytorch/configurations/glm5_next.py index 0982a53f06..fcfd6a2774 100644 --- a/lmdeploy/pytorch/configurations/glm5_next.py +++ b/lmdeploy/pytorch/configurations/glm5_next.py @@ -239,7 +239,7 @@ def build(cls, hf_config, model_path: str | None = None, **kwargs): # DeepSeek-V3.2 token indexer selected by mla_index_topk. Keeping this # unset also preserves the BF16 latent MLA cache policy. config.mla_index_topk = None - config.k_head_dim = text_config.kv_lora_rank + 64 + config.k_head_dim = text_config.kv_lora_rank # Reuse Qwen3.5's token ring for convolution; recurrent/KPool states # keep a complete checkpoint after each verified token. ring_shape = (num_spec_tokens + 1,) if num_spec_tokens else () diff --git a/lmdeploy/pytorch/kernels/cuda/flatten_kv_cache.py b/lmdeploy/pytorch/kernels/cuda/flatten_kv_cache.py index 42ca2b9998..3b9063f833 100644 --- a/lmdeploy/pytorch/kernels/cuda/flatten_kv_cache.py +++ b/lmdeploy/pytorch/kernels/cuda/flatten_kv_cache.py @@ -446,7 +446,7 @@ def flatten_kv_cache(k_caches: Tensor, BLOCK_DK = triton.next_power_of_2(k_head_dim) BLOCK_DV = triton.next_power_of_2(v_head_dim) BLOCK_BS = k_caches.size(s_dim) - shared_kv = k_caches.data_ptr() == v_caches.data_ptr() and v_head_dim < k_head_dim + shared_kv = k_caches.data_ptr() == v_caches.data_ptr() and v_head_dim <= k_head_dim if flatten_kv_layout == 'hsd': k_states = k_caches.new_empty(num_heads, out_size, k_head_dim, dtype=out_dtype) if quant_policy == QuantPolicy.NONE and shared_kv: diff --git a/lmdeploy/pytorch/models/glm5_next.py b/lmdeploy/pytorch/models/glm5_next.py index 88098cbebe..2564679bb2 100644 --- a/lmdeploy/pytorch/models/glm5_next.py +++ b/lmdeploy/pytorch/models/glm5_next.py @@ -742,7 +742,6 @@ class Glm5NextSparseAttention(DeepseekV32Attention): """GLM MLA without RoPE and with a pageable KPool-4 indexer.""" use_sparse_mla = False - mla_head_padding = 64 def __init__(self, config: Any, @@ -1246,10 +1245,6 @@ def forward( (unabsorbed_query, key_states, value_states, q_lora) = self._qkv_proj_unabsorbed( hidden_states, num_heads=num_heads) - # The latent cache retains FlashMLA's 576-wide DeepSeek layout. GLM - # has no RoPE tail, so its final 64 dimensions are exact zeros. - key_states = F.pad(key_states, (0, self.mla_head_padding)) - if not attn_metadata.is_decoding: use_sparse = int(attn_metadata.max_kv_seqlen) > self.index_topk logical_indices = self._kpool_indices( @@ -1301,8 +1296,7 @@ def forward( topk_indices_buffer=topk_indices_buffer, skip_topk=skip_topk, ) - # GLM has no RoPE tail, so the absorbed query contains exactly 512 - # values; the cache retains its 576-wide FlashMLA storage alignment. + # NoPE queries and latent cache both contain exactly 512 values. query_states = self._absorbed_query(unabsorbed_query, num_heads) attn_output = self.decode_attn_fwd.forward( query_states, diff --git a/tests/pytorch/config/test_glm5_runtime_env.py b/tests/pytorch/config/test_glm5_runtime_env.py index 61bf6803ba..e1ad7d5921 100644 --- a/tests/pytorch/config/test_glm5_runtime_env.py +++ b/tests/pytorch/config/test_glm5_runtime_env.py @@ -41,3 +41,30 @@ def test_glm5_leaves_nvls_policy_to_runtime(monkeypatch, nvls): assert seen == [nvls] assert os.environ.get('NCCL_NVLS_ENABLE') == nvls + + +@pytest.mark.parametrize('tp', [1, 4, 8]) +@pytest.mark.parametrize('draft', [False, True]) +def test_glm5_nope_cache_geometry_for_target_and_mtp(monkeypatch, tp, draft): + import torch + + from lmdeploy.pytorch.config import CacheConfig + from lmdeploy.pytorch.engine.cache_engine.schema import build_k_cache_desc, build_v_cache_desc + + hf_config = Glm5NextConfig(text_config={ + 'num_hidden_layers': 4, + 'layer_types': ['linear_attention'] * 3 + ['deepseek_sparse_attention'], + 'linear_num_heads': 64, 'linear_head_dim': 128, 'linear_conv_kernel_dim': 4, + 'index_kpool': 4, 'kv_lora_rank': 512, 'qk_rope_head_dim': 0, + }) + monkeypatch.setattr('lmdeploy.pytorch.configurations.deepseek_v2.flash_mla_available', lambda: True) + config = Glm5NextModelConfigBuilder.build(hf_config, tp=tp, device_type='cuda', + num_spec_tokens=5, is_draft_model=draft) + cache_config = CacheConfig(max_batches=8, block_size=64, num_cpu_blocks=0, num_gpu_blocks=16) + key = build_k_cache_desc(config, cache_config, world_size=tp) + value = build_v_cache_desc(config, cache_config, world_size=tp) + assert key.shape == [64, 1, 512] + assert key.dtype == torch.bfloat16 + assert key.size == 64 * 512 * 2 + assert value.size == 0 + assert not config.use_mla_fp8_cache diff --git a/tests/pytorch/kernel/test_glm5_nope_cache.py b/tests/pytorch/kernel/test_glm5_nope_cache.py new file mode 100644 index 0000000000..110d88c680 --- /dev/null +++ b/tests/pytorch/kernel/test_glm5_nope_cache.py @@ -0,0 +1,103 @@ +# Copyright (c) OpenMMLab. All rights reserved. +from types import SimpleNamespace + +import pytest +import torch +import torch.nn.functional as F + + +@pytest.mark.parametrize('layout', ['hsd', 'shd']) +@pytest.mark.parametrize('storage_width', [512, 576]) +def test_nope_flatten_reuses_shared_value_output(layout, storage_width): + from lmdeploy.pytorch.kernels.cuda.flatten_kv_cache import flatten_kv_cache + + torch.manual_seed(113) + cache = torch.randn(7, 64, 1, storage_width, device='cuda', dtype=torch.bfloat16) + lengths = torch.tensor([67, 19], device='cuda') + blocks = torch.tensor([[3, 1], [5, 2]], device='cuda') + keys, values = flatten_kv_cache(cache, cache[..., :512], lengths, blocks, + out_size=128, flatten_kv_layout=layout) + expected = torch.cat((cache[3], cache[1, :3], cache[5, :19])) + expected = F.pad(expected, (0, 0, 0, 0, 0, 128 - expected.size(0))) + if layout == 'hsd': + expected = expected.transpose(0, 1) + torch.testing.assert_close(keys, expected, rtol=0, atol=0) + torch.testing.assert_close(values, expected[..., :512], rtol=0, atol=0) + assert keys.untyped_storage().data_ptr() == values.untyped_storage().data_ptr() + + +def make_nope_case(lengths, histories, heads, decoding): + from lmdeploy.pytorch.backends.cuda.attention import TritonAttentionMetadata + from lmdeploy.pytorch.backends.cuda.attention.mla import FlashMLAImpl + from lmdeploy.pytorch.backends.cuda.attention.tilelang_sparse_mla import TilelangSparseMLADecode + + torch.manual_seed(127) + q_lens = torch.tensor(lengths, device='cuda') + kv_lens = q_lens + torch.tensor(histories, device='cuda') + q_ends = q_lens.cumsum(0, dtype=torch.int32) + kv_ends = kv_lens.cumsum(0, dtype=torch.int32) + columns = (max(a + b for a, b in zip(lengths, histories)) + 63) // 64 + blocks = torch.randperm(len(lengths) * columns, device='cuda', dtype=torch.int32) + 1 + blocks = blocks.reshape(len(lengths), columns) + metadata = TritonAttentionMetadata( + is_decoding=decoding, block_offsets=blocks, q_start_loc=q_ends - q_lens, + q_seqlens=q_lens, kv_start_loc=kv_ends - kv_lens, kv_seqlens=kv_lens, + cu_seqlens_q=F.pad(q_ends, (1, 0)), cu_seqlens_k=F.pad(kv_ends, (1, 0)), + kv_flatten_size=sum(a + b for a, b in zip(lengths, histories)), + max_kv_seqlen=max(a + b for a, b in zip(lengths, histories)), max_q_seqlen=max(lengths)) + query = torch.randn(sum(lengths), heads, 512, device='cuda', dtype=torch.bfloat16) + key = torch.randn(sum(lengths), 1, 512, device='cuda', dtype=torch.bfloat16) + initial = torch.randn(len(lengths) * columns + 1, 64, 1, 512, device='cuda', dtype=torch.bfloat16) + indices = torch.full((sum(lengths), 2048), -1, device='cuda', dtype=torch.int32) + offset = 0 + for count, history in zip(lengths, histories): + for i in range(count): + seq = history + i + 1 + ids = torch.arange(max(0, seq - 2048), seq, device='cuda', dtype=torch.int32) + indices[offset + i, :ids.numel()] = ids + offset += count + outputs, caches, calls = {}, {}, {} + for width in (576, 512): + cache = F.pad(initial, (0, width - 512)) if width == 576 else initial.clone() + current = F.pad(key, (0, width - 512)) if width == 576 else key + impl = FlashMLAImpl(heads, width, num_kv_heads=1, v_head_size=512) + writer = SimpleNamespace(impl=impl, _lazy_init=lambda device: None, + fill_and_flatten_latent_kv_cache=impl.fill_and_flatten_latent_kv_cache) + backend = TilelangSparseMLADecode(2048, 4) + if decoding: + def call(backend=backend, current=current, cache=cache, writer=writer): + return backend.forward(query, current, current[..., :512], cache, cache[..., :512], + metadata, 192**-0.5, writer, logical_indices=indices) + else: + def call(backend=backend, current=current, cache=cache, writer=writer): + return backend.forward_prefill(query, current, cache, metadata, 192**-0.5, writer, indices) + calls[width] = call + outputs[width] = call() + caches[width] = cache + return outputs, caches, calls, query, metadata + + +@pytest.mark.parametrize('lengths,histories,decoding', [ + ([128], [0], False), ([3, 5], [63, 126], False), ([257, 255], [4095, 8191], False), + ([1, 1], [4095, 127], True), ([6, 6], [4095, 127], True), +]) +@pytest.mark.parametrize('heads', [8, 16]) +def test_nope_512_matches_padded_576_backend(lengths, histories, decoding, heads): + outputs, caches, _, _, _ = make_nope_case(lengths, histories, heads, decoding) + torch.testing.assert_close(outputs[512], outputs[576], rtol=0, atol=0) + torch.testing.assert_close(caches[512], caches[576][..., :512], rtol=0, atol=0) + assert caches[512].nbytes * 9 == caches[576].nbytes * 8 + + +@pytest.mark.parametrize('steps', [1, 6]) +def test_nope_decode_graph_replays_with_new_lengths(steps): + outputs, caches, calls, query, metadata = make_nope_case([steps, steps], [63, 126], 16, True) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + actual = calls[512]() + metadata.kv_seqlens.add_(1) + query.normal_() + expected = calls[576]() + graph.replay() + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + torch.testing.assert_close(caches[512], caches[576][..., :512], rtol=0, atol=0) From 868782a63abab428347fcf053f64da7df45598f6 Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Wed, 23 Sep 2026 17:28:11 +0000 Subject: [PATCH 19/39] perf: read original HC states in shared pre-reduce --- lmdeploy/pytorch/nn/hc_prepost.py | 4 ++-- tests/pytorch/kernel/test_dsv4_hc_prepost.py | 10 ++++++++-- 2 files changed, 10 insertions(+), 4 deletions(-) diff --git a/lmdeploy/pytorch/nn/hc_prepost.py b/lmdeploy/pytorch/nn/hc_prepost.py index 3106b9cb0a..3f7d1d82c8 100644 --- a/lmdeploy/pytorch/nn/hc_prepost.py +++ b/lmdeploy/pytorch/nn/hc_prepost.py @@ -28,7 +28,7 @@ def pre( norm_eps: float, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: from lmdeploy.pytorch.nn.norm import rms_scale - shape, dtype = x.size(), x.dtype + hidden_states, dtype = x, x.dtype x = x.flatten(2).float() if self.avoid_gemv and x.size(0) == 1 and x.size(1) == 1: # Single-token decode otherwise selects GEMV, whose reduction @@ -37,7 +37,7 @@ def pre( else: mixes = F.linear(x, hc_fn) mixes = rms_scale(mixes, x, eps=norm_eps) - return self.impl.pre(x.view(shape), mixes, hc_scale, hc_base, dtype) + return self.impl.pre(hidden_states, mixes, hc_scale, hc_base, dtype) def pre_reduce(self, x: torch.Tensor, pre: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: return self.impl.pre_reduce(x, pre, out_dtype) diff --git a/tests/pytorch/kernel/test_dsv4_hc_prepost.py b/tests/pytorch/kernel/test_dsv4_hc_prepost.py index 08f7db59a0..59ac0c14a1 100644 --- a/tests/pytorch/kernel/test_dsv4_hc_prepost.py +++ b/tests/pytorch/kernel/test_dsv4_hc_prepost.py @@ -40,7 +40,9 @@ def test_pre_reduce(self, lead_shape, dim): assert out.dtype == torch.bfloat16 torch.testing.assert_close(out.float(), ref.float(), atol=1e-2, rtol=1e-2) - def test_pre(self): + @pytest.mark.parametrize('dtype', [torch.bfloat16, torch.float16, torch.float32]) + @pytest.mark.parametrize('expanded', [False, True]) + def test_pre(self, dtype, expanded): from lmdeploy.pytorch.kernels.cuda.dsv4.hc_split_sinkhorn import hc_split_sinkhorn from lmdeploy.pytorch.nn import HcPrePost, rms_scale hc_mult = 4 @@ -52,7 +54,9 @@ def test_pre(self): mix_hc = (2 + hc_mult) * hc_mult hc_dim = hc_mult * dim - x = torch.randn(*lead_shape, hc_mult, dim, device='cuda', dtype=torch.bfloat16) + x = torch.randn(*lead_shape, 1 if expanded else hc_mult, dim, device='cuda', dtype=dtype) + if expanded: + x = x.expand(*lead_shape, hc_mult, dim) hc_fn = torch.randn(mix_hc, hc_dim, device='cuda', dtype=torch.float32) hc_scale = torch.randn(3, device='cuda', dtype=torch.float32) hc_base = torch.randn(mix_hc, device='cuda', dtype=torch.float32) @@ -65,10 +69,12 @@ def test_pre(self): pre_ref, post_ref, comb_ref = hc_split_sinkhorn( mixes, hc_scale, hc_base, hc_mult, sinkhorn_iters, sinkhorn_eps) out_ref = _reference_pre_reduce(x_flat.view_as(x), pre_ref, x.dtype) + fp32_out = op.pre_reduce(x_flat.view_as(x), pre_ref, x.dtype) assert out.shape == (*lead_shape, dim) assert post.shape == (*lead_shape, hc_mult) assert comb.shape == (*lead_shape, hc_mult, hc_mult) + torch.testing.assert_close(out, fp32_out, atol=0, rtol=0) torch.testing.assert_close(out.float(), out_ref.float(), atol=1e-2, rtol=1e-2) torch.testing.assert_close(post, post_ref, atol=1e-6, rtol=1e-6) torch.testing.assert_close(comb, comb_ref, atol=1e-6, rtol=1e-6) From 5988cd5cfef76e382db6b944b4bb7f2fa5a13606 Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Wed, 23 Sep 2026 17:48:32 +0000 Subject: [PATCH 20/39] perf: use compact FP8 MoE scheduling for sparse routes --- .../pytorch/kernels/cuda/moe/blocked_fp8.py | 5 +++ .../kernel/test_fuse_moe_blocked_fp8.py | 42 +++++++++++++++++++ 2 files changed, 47 insertions(+) diff --git a/lmdeploy/pytorch/kernels/cuda/moe/blocked_fp8.py b/lmdeploy/pytorch/kernels/cuda/moe/blocked_fp8.py index b12c69b2ab..b9dc343fb6 100644 --- a/lmdeploy/pytorch/kernels/cuda/moe/blocked_fp8.py +++ b/lmdeploy/pytorch/kernels/cuda/moe/blocked_fp8.py @@ -633,6 +633,11 @@ def _select_compact_blocked_fp8_moe_both_config(num_tokens: int, num_routes: int avg_routes = triton.cdiv(num_routes, num_experts) if local_experts != num_experts or local_experts < 256: return None + if (avg_routes == 1 and local_experts == 288 + and gate_out_features == 1024 and input_features == 4096): + # Sparse routing at this width benefits from skipping inactive + # experts while retaining the small-M reduction schedule. + return dict(block_m=16, block_n=128, num_warps=4, num_stages=3) n_tiles = gate_out_features // 128 k_blocks = input_features // 128 if n_tiles >= 8: diff --git a/tests/pytorch/kernel/test_fuse_moe_blocked_fp8.py b/tests/pytorch/kernel/test_fuse_moe_blocked_fp8.py index 587b930df8..e360fed0f4 100644 --- a/tests/pytorch/kernel/test_fuse_moe_blocked_fp8.py +++ b/tests/pytorch/kernel/test_fuse_moe_blocked_fp8.py @@ -76,6 +76,48 @@ def test_compact_blocked_fp8_gate_config(block_m, block_n, input_features, expec assert gate_config.get('transpose_mma', False) is transpose_mma +@pytest.mark.skipif(torch.cuda.get_device_capability()[0] < 9, reason='require device with cc>=9.0') +@pytest.mark.parametrize('tokens', [1, 6, 36, 37]) +@pytest.mark.parametrize('concentrated', [False, True]) +@torch.inference_mode() +def test_sparse_route_blocked_fp8_preserves_fp32_reduction(monkeypatch, tokens, concentrated): + import importlib + from functools import partial + + from lmdeploy.pytorch.kernels.cuda.activation import silu_and_mul + from lmdeploy.pytorch.kernels.cuda.blocked_gemm_fp8 import quant_fp8 + + module = importlib.import_module('lmdeploy.pytorch.kernels.cuda.moe.blocked_fp8') + torch.manual_seed(33) + experts, hidden, intermediate, topk = 288, 4096, 512, 8 + dtype = torch.float8_e4m3fn + w1 = torch.randint(-4, 5, (experts, 2 * intermediate, hidden), device='cuda', dtype=torch.int8).to(dtype) + w2 = torch.randint(-4, 5, (experts, hidden, intermediate), device='cuda', dtype=torch.int8).to(dtype) + s1 = torch.rand(experts, 2 * intermediate // 128, hidden // 128, device='cuda') * .01 + .001 + s2 = torch.rand(experts, hidden // 128, intermediate // 128, device='cuda') * .01 + .001 + x = torch.randn(tokens, hidden, device='cuda', dtype=torch.bfloat16) * .1 + scores = torch.randn(tokens, experts, device='cuda') + if concentrated: + scores[:, :-topk] = -float('inf') + weights, ids = scores.topk(topk, dim=-1) + weights = weights.softmax(-1) + quant, scales = quant_fp8(x, 128, dtype=dtype) + + def run(): + return module.fused_moe_blocked_fp8( + quant, scales, w1, s1, w2, s2, weights, ids, topk, + out_dtype=torch.bfloat16, fp32_acc=True, output_scale=2.5, + act_func=partial(silu_and_mul, swiglu_limit=10., precise_mul=True)) + + select = module._select_compact_blocked_fp8_moe_both_config + with monkeypatch.context() as patch: + patch.setattr(module, '_select_compact_blocked_fp8_moe_both_config', + lambda *args: None if args[1] <= args[2] else select(*args)) + expected = run() + actual = run() + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + + @pytest.mark.skipif(torch.cuda.get_device_capability()[0] < 9, reason='require device with cc>=9.0') @torch.inference_mode() def test_fused_moe_blocked_fp8_compact_transposed_mma_matches_normal(): From 82e451c8628e967eefc75db67bf19e9d2a76ee56 Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Thu, 24 Sep 2026 02:58:25 +0000 Subject: [PATCH 21/39] revert: remove PR-specific test changes --- tests/pytorch/config/test_glm5_runtime_env.py | 70 -------- tests/pytorch/kernel/test_dsv4_hc_prepost.py | 10 +- .../kernel/test_fuse_moe_blocked_fp8.py | 42 ----- tests/pytorch/kernel/test_glm5_nope_cache.py | 103 ----------- tests/pytorch/kernel/test_kda_fused_gate.py | 120 ------------- tests/pytorch/kernel/test_kda_strided.py | 67 ------- tests/pytorch/kernel/test_kpool_prefill.py | 164 ------------------ tests/pytorch/kernel/test_kpool_selection.py | 77 -------- tests/pytorch/nn/test_glm5_layer_norm.py | 67 ------- tests/pytorch/nn/test_moe_options.py | 122 ------------- tests/pytorch/nn/test_rotary_embedding.py | 67 ------- 11 files changed, 2 insertions(+), 907 deletions(-) delete mode 100644 tests/pytorch/config/test_glm5_runtime_env.py delete mode 100644 tests/pytorch/kernel/test_glm5_nope_cache.py delete mode 100644 tests/pytorch/kernel/test_kda_fused_gate.py delete mode 100644 tests/pytorch/kernel/test_kda_strided.py delete mode 100644 tests/pytorch/kernel/test_kpool_prefill.py delete mode 100644 tests/pytorch/kernel/test_kpool_selection.py delete mode 100644 tests/pytorch/nn/test_glm5_layer_norm.py delete mode 100644 tests/pytorch/nn/test_moe_options.py diff --git a/tests/pytorch/config/test_glm5_runtime_env.py b/tests/pytorch/config/test_glm5_runtime_env.py deleted file mode 100644 index e1ad7d5921..0000000000 --- a/tests/pytorch/config/test_glm5_runtime_env.py +++ /dev/null @@ -1,70 +0,0 @@ -# Copyright (c) OpenMMLab. All rights reserved. -import os -from unittest.mock import Mock - -import pytest - -from lmdeploy.hf_configs.configuration_glm5_next import Glm5NextConfig -from lmdeploy.pytorch.config import DistConfig -from lmdeploy.pytorch.configurations.glm5_next import Glm5NextModelConfigBuilder -from lmdeploy.pytorch.engine.executor import base_worker - - -@pytest.mark.parametrize('nvls', [None, '0', '1']) -def test_glm5_leaves_nvls_policy_to_runtime(monkeypatch, nvls): - if nvls is None: - monkeypatch.delenv('NCCL_NVLS_ENABLE', raising=False) - else: - monkeypatch.setenv('NCCL_NVLS_ENABLE', nvls) - hf_config = Glm5NextConfig(text_config={ - 'num_hidden_layers': 4, - 'layer_types': ['linear_attention'] * 3 + ['deepseek_sparse_attention'], - 'linear_num_heads': 64, - 'linear_head_dim': 128, - 'linear_conv_kernel_dim': 4, - 'index_kpool': 4, - }) - monkeypatch.setattr('lmdeploy.pytorch.configurations.deepseek_v2.flash_mla_available', lambda: True) - config = Glm5NextModelConfigBuilder.build(hf_config, tp=8, device_type='cuda') - worker = base_worker.WorkerWrapperBase.__new__(base_worker.WorkerWrapperBase) - worker.model_config = config - worker.dist_config = DistConfig(tp=8) - worker.world_size = 8 - worker.device_type = 'cuda' - seen = [] - monkeypatch.setattr(base_worker, 'init_process_group', - lambda rank, size: seen.append(os.environ.get('NCCL_NVLS_ENABLE'))) - monkeypatch.setattr(base_worker, 'get_backend', Mock()) - monkeypatch.setattr(base_worker.DistContext, 'build', Mock()) - - worker.init_process_group(rank=0) - - assert seen == [nvls] - assert os.environ.get('NCCL_NVLS_ENABLE') == nvls - - -@pytest.mark.parametrize('tp', [1, 4, 8]) -@pytest.mark.parametrize('draft', [False, True]) -def test_glm5_nope_cache_geometry_for_target_and_mtp(monkeypatch, tp, draft): - import torch - - from lmdeploy.pytorch.config import CacheConfig - from lmdeploy.pytorch.engine.cache_engine.schema import build_k_cache_desc, build_v_cache_desc - - hf_config = Glm5NextConfig(text_config={ - 'num_hidden_layers': 4, - 'layer_types': ['linear_attention'] * 3 + ['deepseek_sparse_attention'], - 'linear_num_heads': 64, 'linear_head_dim': 128, 'linear_conv_kernel_dim': 4, - 'index_kpool': 4, 'kv_lora_rank': 512, 'qk_rope_head_dim': 0, - }) - monkeypatch.setattr('lmdeploy.pytorch.configurations.deepseek_v2.flash_mla_available', lambda: True) - config = Glm5NextModelConfigBuilder.build(hf_config, tp=tp, device_type='cuda', - num_spec_tokens=5, is_draft_model=draft) - cache_config = CacheConfig(max_batches=8, block_size=64, num_cpu_blocks=0, num_gpu_blocks=16) - key = build_k_cache_desc(config, cache_config, world_size=tp) - value = build_v_cache_desc(config, cache_config, world_size=tp) - assert key.shape == [64, 1, 512] - assert key.dtype == torch.bfloat16 - assert key.size == 64 * 512 * 2 - assert value.size == 0 - assert not config.use_mla_fp8_cache diff --git a/tests/pytorch/kernel/test_dsv4_hc_prepost.py b/tests/pytorch/kernel/test_dsv4_hc_prepost.py index 59ac0c14a1..08f7db59a0 100644 --- a/tests/pytorch/kernel/test_dsv4_hc_prepost.py +++ b/tests/pytorch/kernel/test_dsv4_hc_prepost.py @@ -40,9 +40,7 @@ def test_pre_reduce(self, lead_shape, dim): assert out.dtype == torch.bfloat16 torch.testing.assert_close(out.float(), ref.float(), atol=1e-2, rtol=1e-2) - @pytest.mark.parametrize('dtype', [torch.bfloat16, torch.float16, torch.float32]) - @pytest.mark.parametrize('expanded', [False, True]) - def test_pre(self, dtype, expanded): + def test_pre(self): from lmdeploy.pytorch.kernels.cuda.dsv4.hc_split_sinkhorn import hc_split_sinkhorn from lmdeploy.pytorch.nn import HcPrePost, rms_scale hc_mult = 4 @@ -54,9 +52,7 @@ def test_pre(self, dtype, expanded): mix_hc = (2 + hc_mult) * hc_mult hc_dim = hc_mult * dim - x = torch.randn(*lead_shape, 1 if expanded else hc_mult, dim, device='cuda', dtype=dtype) - if expanded: - x = x.expand(*lead_shape, hc_mult, dim) + x = torch.randn(*lead_shape, hc_mult, dim, device='cuda', dtype=torch.bfloat16) hc_fn = torch.randn(mix_hc, hc_dim, device='cuda', dtype=torch.float32) hc_scale = torch.randn(3, device='cuda', dtype=torch.float32) hc_base = torch.randn(mix_hc, device='cuda', dtype=torch.float32) @@ -69,12 +65,10 @@ def test_pre(self, dtype, expanded): pre_ref, post_ref, comb_ref = hc_split_sinkhorn( mixes, hc_scale, hc_base, hc_mult, sinkhorn_iters, sinkhorn_eps) out_ref = _reference_pre_reduce(x_flat.view_as(x), pre_ref, x.dtype) - fp32_out = op.pre_reduce(x_flat.view_as(x), pre_ref, x.dtype) assert out.shape == (*lead_shape, dim) assert post.shape == (*lead_shape, hc_mult) assert comb.shape == (*lead_shape, hc_mult, hc_mult) - torch.testing.assert_close(out, fp32_out, atol=0, rtol=0) torch.testing.assert_close(out.float(), out_ref.float(), atol=1e-2, rtol=1e-2) torch.testing.assert_close(post, post_ref, atol=1e-6, rtol=1e-6) torch.testing.assert_close(comb, comb_ref, atol=1e-6, rtol=1e-6) diff --git a/tests/pytorch/kernel/test_fuse_moe_blocked_fp8.py b/tests/pytorch/kernel/test_fuse_moe_blocked_fp8.py index e360fed0f4..587b930df8 100644 --- a/tests/pytorch/kernel/test_fuse_moe_blocked_fp8.py +++ b/tests/pytorch/kernel/test_fuse_moe_blocked_fp8.py @@ -76,48 +76,6 @@ def test_compact_blocked_fp8_gate_config(block_m, block_n, input_features, expec assert gate_config.get('transpose_mma', False) is transpose_mma -@pytest.mark.skipif(torch.cuda.get_device_capability()[0] < 9, reason='require device with cc>=9.0') -@pytest.mark.parametrize('tokens', [1, 6, 36, 37]) -@pytest.mark.parametrize('concentrated', [False, True]) -@torch.inference_mode() -def test_sparse_route_blocked_fp8_preserves_fp32_reduction(monkeypatch, tokens, concentrated): - import importlib - from functools import partial - - from lmdeploy.pytorch.kernels.cuda.activation import silu_and_mul - from lmdeploy.pytorch.kernels.cuda.blocked_gemm_fp8 import quant_fp8 - - module = importlib.import_module('lmdeploy.pytorch.kernels.cuda.moe.blocked_fp8') - torch.manual_seed(33) - experts, hidden, intermediate, topk = 288, 4096, 512, 8 - dtype = torch.float8_e4m3fn - w1 = torch.randint(-4, 5, (experts, 2 * intermediate, hidden), device='cuda', dtype=torch.int8).to(dtype) - w2 = torch.randint(-4, 5, (experts, hidden, intermediate), device='cuda', dtype=torch.int8).to(dtype) - s1 = torch.rand(experts, 2 * intermediate // 128, hidden // 128, device='cuda') * .01 + .001 - s2 = torch.rand(experts, hidden // 128, intermediate // 128, device='cuda') * .01 + .001 - x = torch.randn(tokens, hidden, device='cuda', dtype=torch.bfloat16) * .1 - scores = torch.randn(tokens, experts, device='cuda') - if concentrated: - scores[:, :-topk] = -float('inf') - weights, ids = scores.topk(topk, dim=-1) - weights = weights.softmax(-1) - quant, scales = quant_fp8(x, 128, dtype=dtype) - - def run(): - return module.fused_moe_blocked_fp8( - quant, scales, w1, s1, w2, s2, weights, ids, topk, - out_dtype=torch.bfloat16, fp32_acc=True, output_scale=2.5, - act_func=partial(silu_and_mul, swiglu_limit=10., precise_mul=True)) - - select = module._select_compact_blocked_fp8_moe_both_config - with monkeypatch.context() as patch: - patch.setattr(module, '_select_compact_blocked_fp8_moe_both_config', - lambda *args: None if args[1] <= args[2] else select(*args)) - expected = run() - actual = run() - torch.testing.assert_close(actual, expected, rtol=0, atol=0) - - @pytest.mark.skipif(torch.cuda.get_device_capability()[0] < 9, reason='require device with cc>=9.0') @torch.inference_mode() def test_fused_moe_blocked_fp8_compact_transposed_mma_matches_normal(): diff --git a/tests/pytorch/kernel/test_glm5_nope_cache.py b/tests/pytorch/kernel/test_glm5_nope_cache.py deleted file mode 100644 index 110d88c680..0000000000 --- a/tests/pytorch/kernel/test_glm5_nope_cache.py +++ /dev/null @@ -1,103 +0,0 @@ -# Copyright (c) OpenMMLab. All rights reserved. -from types import SimpleNamespace - -import pytest -import torch -import torch.nn.functional as F - - -@pytest.mark.parametrize('layout', ['hsd', 'shd']) -@pytest.mark.parametrize('storage_width', [512, 576]) -def test_nope_flatten_reuses_shared_value_output(layout, storage_width): - from lmdeploy.pytorch.kernels.cuda.flatten_kv_cache import flatten_kv_cache - - torch.manual_seed(113) - cache = torch.randn(7, 64, 1, storage_width, device='cuda', dtype=torch.bfloat16) - lengths = torch.tensor([67, 19], device='cuda') - blocks = torch.tensor([[3, 1], [5, 2]], device='cuda') - keys, values = flatten_kv_cache(cache, cache[..., :512], lengths, blocks, - out_size=128, flatten_kv_layout=layout) - expected = torch.cat((cache[3], cache[1, :3], cache[5, :19])) - expected = F.pad(expected, (0, 0, 0, 0, 0, 128 - expected.size(0))) - if layout == 'hsd': - expected = expected.transpose(0, 1) - torch.testing.assert_close(keys, expected, rtol=0, atol=0) - torch.testing.assert_close(values, expected[..., :512], rtol=0, atol=0) - assert keys.untyped_storage().data_ptr() == values.untyped_storage().data_ptr() - - -def make_nope_case(lengths, histories, heads, decoding): - from lmdeploy.pytorch.backends.cuda.attention import TritonAttentionMetadata - from lmdeploy.pytorch.backends.cuda.attention.mla import FlashMLAImpl - from lmdeploy.pytorch.backends.cuda.attention.tilelang_sparse_mla import TilelangSparseMLADecode - - torch.manual_seed(127) - q_lens = torch.tensor(lengths, device='cuda') - kv_lens = q_lens + torch.tensor(histories, device='cuda') - q_ends = q_lens.cumsum(0, dtype=torch.int32) - kv_ends = kv_lens.cumsum(0, dtype=torch.int32) - columns = (max(a + b for a, b in zip(lengths, histories)) + 63) // 64 - blocks = torch.randperm(len(lengths) * columns, device='cuda', dtype=torch.int32) + 1 - blocks = blocks.reshape(len(lengths), columns) - metadata = TritonAttentionMetadata( - is_decoding=decoding, block_offsets=blocks, q_start_loc=q_ends - q_lens, - q_seqlens=q_lens, kv_start_loc=kv_ends - kv_lens, kv_seqlens=kv_lens, - cu_seqlens_q=F.pad(q_ends, (1, 0)), cu_seqlens_k=F.pad(kv_ends, (1, 0)), - kv_flatten_size=sum(a + b for a, b in zip(lengths, histories)), - max_kv_seqlen=max(a + b for a, b in zip(lengths, histories)), max_q_seqlen=max(lengths)) - query = torch.randn(sum(lengths), heads, 512, device='cuda', dtype=torch.bfloat16) - key = torch.randn(sum(lengths), 1, 512, device='cuda', dtype=torch.bfloat16) - initial = torch.randn(len(lengths) * columns + 1, 64, 1, 512, device='cuda', dtype=torch.bfloat16) - indices = torch.full((sum(lengths), 2048), -1, device='cuda', dtype=torch.int32) - offset = 0 - for count, history in zip(lengths, histories): - for i in range(count): - seq = history + i + 1 - ids = torch.arange(max(0, seq - 2048), seq, device='cuda', dtype=torch.int32) - indices[offset + i, :ids.numel()] = ids - offset += count - outputs, caches, calls = {}, {}, {} - for width in (576, 512): - cache = F.pad(initial, (0, width - 512)) if width == 576 else initial.clone() - current = F.pad(key, (0, width - 512)) if width == 576 else key - impl = FlashMLAImpl(heads, width, num_kv_heads=1, v_head_size=512) - writer = SimpleNamespace(impl=impl, _lazy_init=lambda device: None, - fill_and_flatten_latent_kv_cache=impl.fill_and_flatten_latent_kv_cache) - backend = TilelangSparseMLADecode(2048, 4) - if decoding: - def call(backend=backend, current=current, cache=cache, writer=writer): - return backend.forward(query, current, current[..., :512], cache, cache[..., :512], - metadata, 192**-0.5, writer, logical_indices=indices) - else: - def call(backend=backend, current=current, cache=cache, writer=writer): - return backend.forward_prefill(query, current, cache, metadata, 192**-0.5, writer, indices) - calls[width] = call - outputs[width] = call() - caches[width] = cache - return outputs, caches, calls, query, metadata - - -@pytest.mark.parametrize('lengths,histories,decoding', [ - ([128], [0], False), ([3, 5], [63, 126], False), ([257, 255], [4095, 8191], False), - ([1, 1], [4095, 127], True), ([6, 6], [4095, 127], True), -]) -@pytest.mark.parametrize('heads', [8, 16]) -def test_nope_512_matches_padded_576_backend(lengths, histories, decoding, heads): - outputs, caches, _, _, _ = make_nope_case(lengths, histories, heads, decoding) - torch.testing.assert_close(outputs[512], outputs[576], rtol=0, atol=0) - torch.testing.assert_close(caches[512], caches[576][..., :512], rtol=0, atol=0) - assert caches[512].nbytes * 9 == caches[576].nbytes * 8 - - -@pytest.mark.parametrize('steps', [1, 6]) -def test_nope_decode_graph_replays_with_new_lengths(steps): - outputs, caches, calls, query, metadata = make_nope_case([steps, steps], [63, 126], 16, True) - graph = torch.cuda.CUDAGraph() - with torch.cuda.graph(graph): - actual = calls[512]() - metadata.kv_seqlens.add_(1) - query.normal_() - expected = calls[576]() - graph.replay() - torch.testing.assert_close(actual, expected, rtol=0, atol=0) - torch.testing.assert_close(caches[512], caches[576][..., :512], rtol=0, atol=0) diff --git a/tests/pytorch/kernel/test_kda_fused_gate.py b/tests/pytorch/kernel/test_kda_fused_gate.py deleted file mode 100644 index 2b035402a2..0000000000 --- a/tests/pytorch/kernel/test_kda_fused_gate.py +++ /dev/null @@ -1,120 +0,0 @@ -# Copyright (c) OpenMMLab. All rights reserved. -import pytest -import torch - - -@pytest.mark.parametrize('steps', [1, 3, 6]) -@pytest.mark.parametrize('state_dtype', [torch.float32, torch.bfloat16]) -@pytest.mark.parametrize('with_bias', [False, True]) -def test_fused_kda_gate_preserves_recurrence_and_dummy_state(steps, state_dtype, with_bias): - from fla.ops.kda.gate import kda_gate_fwd - - from lmdeploy.pytorch.kernels.cuda.gated_delta_rule import fused_recurrent_gated_delta_rule as run - - torch.manual_seed(7) - batch, heads, dim, ring = 3, 2, 128, 6 - mixed = torch.randn(batch, steps, 3 * heads * dim, device='cuda', dtype=torch.bfloat16) - q, k, v = [x.unflatten(-1, (heads, dim)) for x in mixed.chunk(3, dim=-1)] - raw_gate = torch.randn_like(q).contiguous() - raw_beta = torch.randn(batch, steps, heads, device='cuda', dtype=torch.bfloat16) - a_log = torch.randn(heads, device='cuda') - dt_bias = torch.randn(heads * dim, device='cuda') if with_bias else None - initial = torch.randn(batch, ring, heads, dim, dim, device='cuda', dtype=state_dtype) * 0.1 - state = initial.clone() - kwargs = dict(state_indices=torch.tensor([2, -1, 0], device='cuda'), - cache_seqlens=torch.tensor([5, 0, 3], device='cuda', dtype=torch.int32), - output_final_state=True, transpose_state_layout=True, use_qk_l2norm_in_kernel=True) - gate = kda_gate_fwd(raw_gate, a_log, dt_bias, lower_bound=-5.0) - expected, expected_state = run(q, k, v, g=gate, beta=raw_beta.float().sigmoid(), - initial_state=initial.clone(), **kwargs) - actual, _ = run(q, k, v, g=raw_gate, beta=raw_beta, a_log=a_log, dt_bias=dt_bias, - lower_bound=-5.0, initial_state=state, **kwargs) - torch.testing.assert_close(actual, expected, rtol=0, atol=0) - torch.testing.assert_close(state, expected_state, rtol=0, atol=0) - torch.testing.assert_close(state[1], initial[1], rtol=0, atol=0) - assert torch.count_nonzero(actual[1]) == 0 - - -def test_fused_kda_gate_cuda_graph_replay(): - from lmdeploy.pytorch.kernels.cuda.gated_delta_rule import fused_recurrent_gated_delta_rule as run - - torch.manual_seed(29) - q, k, v, raw_gate = [torch.randn(1, 6, 2, 128, device='cuda', dtype=torch.bfloat16) for _ in range(4)] - raw_beta = torch.randn(1, 6, 2, device='cuda', dtype=torch.bfloat16) - initial = torch.randn(1, 6, 2, 128, 128, device='cuda') - state = initial.clone() - history = torch.tensor([5], device='cuda', dtype=torch.int32) - kwargs = dict(g=raw_gate, beta=raw_beta, a_log=torch.randn(2, device='cuda'), - dt_bias=torch.randn(256, device='cuda'), lower_bound=-5.0, cache_seqlens=history, - output_final_state=True, transpose_state_layout=True, use_qk_l2norm_in_kernel=True) - run(q, k, v, initial_state=state, **kwargs) - graph = torch.cuda.CUDAGraph() - with torch.cuda.graph(graph): - actual, _ = run(q, k, v, initial_state=state, **kwargs) - raw_gate.normal_() - raw_beta.normal_() - history.fill_(2) - expected, expected_state = run(q, k, v, initial_state=initial.clone(), **kwargs) - state.copy_(initial) - graph.replay() - torch.testing.assert_close(actual, expected, rtol=0, atol=0) - torch.testing.assert_close(state, expected_state, rtol=0, atol=0) - - -def test_kda_unbounded_gate_keeps_existing_arithmetic(): - from lmdeploy.pytorch.backends.cuda.kda import CudaKdaImpl - - torch.manual_seed(43) - impl = CudaKdaImpl() - q, k, v, raw_gate = [torch.randn(1, 1, 2, 128, device='cuda', dtype=torch.bfloat16) for _ in range(4)] - raw_beta = torch.randn(1, 1, 2, device='cuda', dtype=torch.bfloat16) - a_log = torch.randn(2, device='cuda') - dt_bias = torch.randn(256, device='cuda') - initial = torch.randn(1, 2, 128, 128, device='cuda') - gate = impl.kda_gate(raw_gate, a_log, dt_bias) - expected, expected_state = impl.recurrent_func( - q, k, v, g=gate, beta=raw_beta.float().sigmoid(), initial_state=initial.clone(), - output_final_state=True, transpose_state_layout=True, use_qk_l2norm_in_kernel=True) - actual, state = impl._decode_recurrent( - q, k, v, raw_gate, raw_beta, a_log, dt_bias, initial.clone(), lower_bound=None) - torch.testing.assert_close(actual, expected, rtol=0, atol=0) - torch.testing.assert_close(state, expected_state, rtol=0, atol=0) - - -@pytest.mark.parametrize('beta_dtype', [torch.bfloat16, torch.float32]) -def test_fused_kda_gate_extreme_beta(beta_dtype): - from fla.ops.kda.gate import kda_gate_fwd - - from lmdeploy.pytorch.kernels.cuda.gated_delta_rule import fused_recurrent_gated_delta_rule as run - - torch.manual_seed(53) - q, k, v, raw_gate = [torch.randn(9, 1, 2, 128, device='cuda', dtype=torch.bfloat16) for _ in range(4)] - raw_beta = torch.tensor([-100, -89, -88, -87.5, -86, 0, 86, 88, 100], device='cuda', dtype=beta_dtype) - raw_beta = raw_beta[:, None, None].expand(9, 1, 2).contiguous() - a_log = torch.randn(2, device='cuda') - initial = torch.randn(9, 2, 128, 128, device='cuda') - kwargs = dict(output_final_state=True, transpose_state_layout=True, use_qk_l2norm_in_kernel=True) - expected, expected_state = run( - q, k, v, g=kda_gate_fwd(raw_gate, a_log, lower_bound=-5.0), beta=raw_beta.float().sigmoid(), - initial_state=initial.clone(), **kwargs) - actual, state = run(q, k, v, g=raw_gate, beta=raw_beta, a_log=a_log, lower_bound=-5.0, - initial_state=initial.clone(), **kwargs) - torch.testing.assert_close(actual, expected, rtol=0, atol=0) - torch.testing.assert_close(state, expected_state, rtol=0, atol=0) - - -@pytest.mark.parametrize('invalid', [ - 'missing_bound', 'positive_bound', 'missing_beta', 'untransposed', 'missing_a_log' -]) -def test_fused_kda_gate_rejects_unsupported_contract(invalid): - from lmdeploy.pytorch.kernels.cuda.gated_delta_rule import fused_recurrent_gated_delta_rule as run - - q = torch.empty(1, 1, 2, 128) - kwargs = dict(g=q, beta=torch.empty(1, 1, 2), a_log=torch.zeros(2), - dt_bias=torch.zeros(256), lower_bound=-5.0, transpose_state_layout=True) - overrides = dict(missing_bound=dict(lower_bound=None), positive_bound=dict(lower_bound=1.0), - missing_beta=dict(beta=None), untransposed=dict(transpose_state_layout=False), - missing_a_log=dict(a_log=None)) - kwargs.update(overrides[invalid]) - with pytest.raises(ValueError, match='KDA gating'): - run(q, q, q, initial_state=torch.empty(1, 2, 128, 128), **kwargs) diff --git a/tests/pytorch/kernel/test_kda_strided.py b/tests/pytorch/kernel/test_kda_strided.py deleted file mode 100644 index 5889fec4d8..0000000000 --- a/tests/pytorch/kernel/test_kda_strided.py +++ /dev/null @@ -1,67 +0,0 @@ -# Copyright (c) OpenMMLab. All rights reserved. -import pytest -import torch - - -@pytest.mark.parametrize('batch', [1, 3]) -@pytest.mark.parametrize('steps', [1, 3, 6]) -@pytest.mark.parametrize('channel_major', [False, True]) -@pytest.mark.parametrize('state_dtype', [torch.float32, torch.bfloat16]) -def test_kda_strided_inputs_match_contiguous_state_ring(batch, steps, channel_major, state_dtype): - from lmdeploy.pytorch.kernels.cuda.gated_delta_rule import fused_recurrent_gated_delta_rule - - torch.manual_seed(17) - heads, dim, ring = 2, 128, 6 - if channel_major: - mixed = torch.randn(batch, 3 * heads * dim, steps, device='cuda', dtype=torch.bfloat16).transpose(1, 2) - else: - mixed = torch.randn(batch, steps, 3 * heads * dim, device='cuda', dtype=torch.bfloat16) - q, k, v = [x.unflatten(-1, (heads, dim)) for x in mixed.chunk(3, dim=-1)] - gate = -torch.rand(batch, steps, heads, dim, device='cuda') - beta = torch.rand(batch, steps, heads, device='cuda') - initial = torch.randn(3, ring, heads, dim, dim, device='cuda', dtype=state_dtype) * 0.1 - ids = torch.tensor([2, -1, 0][:batch], device='cuda') - history = torch.tensor([5, 0, 3][:batch], device='cuda', dtype=torch.int32) - kwargs = dict(g=gate, beta=beta, state_indices=ids, cache_seqlens=history, - output_final_state=True, transpose_state_layout=True, use_qk_l2norm_in_kernel=True) - expected, expected_state = fused_recurrent_gated_delta_rule( - q.contiguous(), k.contiguous(), v.contiguous(), initial_state=initial.clone(), **kwargs) - actual, actual_state = fused_recurrent_gated_delta_rule(q, k, v, initial_state=initial.clone(), **kwargs) - assert actual.is_contiguous() - torch.testing.assert_close(actual, expected, rtol=0, atol=0) - torch.testing.assert_close(actual_state, expected_state, rtol=0, atol=0) - torch.testing.assert_close(actual_state[1], initial[1], rtol=0, atol=0) - if batch == 3: - assert torch.count_nonzero(actual[1]) == 0 - - -def test_kda_strided_graph_replay_changes_inputs_and_history(): - from lmdeploy.pytorch.kernels.cuda.gated_delta_rule import fused_recurrent_gated_delta_rule - - torch.manual_seed(31) - heads, dim, steps = 2, 128, 6 - mixed = torch.randn(1, 3 * heads * dim, steps, device='cuda', dtype=torch.bfloat16) - q, k, v = [x.unflatten(-1, (heads, dim)) for x in mixed.transpose(1, 2).chunk(3, dim=-1)] - gate = -torch.rand_like(q, dtype=torch.float32).contiguous() - beta = torch.rand(1, steps, heads, device='cuda') - initial = torch.randn(1, steps, heads, dim, dim, device='cuda') * 0.1 - state = initial.clone() - history = torch.tensor([5], device='cuda', dtype=torch.int32) - kwargs = dict(g=gate, beta=beta, cache_seqlens=history, output_final_state=True, - transpose_state_layout=True, use_qk_l2norm_in_kernel=True) - - def run(): - return fused_recurrent_gated_delta_rule(q, k, v, initial_state=state, **kwargs) - - run() - graph = torch.cuda.CUDAGraph() - with torch.cuda.graph(graph): - actual, _ = run() - mixed.normal_() - history.fill_(2) - state.copy_(initial) - expected, expected_state = fused_recurrent_gated_delta_rule( - q.contiguous(), k.contiguous(), v.contiguous(), initial_state=initial.clone(), **kwargs) - graph.replay() - torch.testing.assert_close(actual, expected, rtol=0, atol=0) - torch.testing.assert_close(state, expected_state, rtol=0, atol=0) diff --git a/tests/pytorch/kernel/test_kpool_prefill.py b/tests/pytorch/kernel/test_kpool_prefill.py deleted file mode 100644 index c060ce787f..0000000000 --- a/tests/pytorch/kernel/test_kpool_prefill.py +++ /dev/null @@ -1,164 +0,0 @@ -# Copyright (c) OpenMMLab. All rights reserved. -import pytest -import torch - - -def reference_update(keys, scores, states, ids, q_lens, kv_lens, cache, blocks, ape, round_scale): - from lmdeploy.pytorch.backends.cuda.kpool import kpool_compress_quantize_cuda - from lmdeploy.pytorch.nn.kpool import kpool_partition_update, kpool_write_packed_cache - - start = 0 - for request, (q_len, kv_len, state_id) in enumerate(zip(q_lens.tolist(), kv_lens.tolist(), ids.tolist())): - if state_id >= 0: - history = kv_len - q_len - old_tail = history % 4 - update = kpool_partition_update(keys[start:start + q_len], scores[start:start + q_len], history, 4, - states[0][state_id, :old_tail], states[1][state_id, :old_tail]) - if update.closed_group_ids.numel(): - values, scales = kpool_compress_quantize_cuda( - update.closed_keys, update.closed_scores, ape, mode='extend', round_scale=round_scale) - kpool_write_packed_cache(cache, blocks[request], update.closed_group_ids, values, scales, 4) - for state, tail in zip(states, (update.tail_keys, update.tail_scores)): - # Clone because the empty-query reference can alias the source tail. - tail = tail.clone() - state[state_id].zero_() - state[state_id, :tail.size(0)].copy_(tail) - start += q_len - - -def make_case(lengths, histories, round_scale): - torch.manual_seed(17) - batch = len(lengths) - keys = torch.randn(sum(lengths), 128, device='cuda', dtype=torch.bfloat16) - scores = torch.randn_like(keys) - q_lens = torch.tensor(lengths, device='cuda', dtype=torch.int32) - kv_lens = q_lens + torch.tensor(histories, device='cuda', dtype=torch.int32) - ids = torch.arange(batch, 0, -1, device='cuda') - states = tuple(torch.randn(batch + 2, 4, 128, device='cuda', dtype=torch.bfloat16) for _ in range(2)) - blocks = torch.arange(1, batch * 160 + 1, device='cuda').reshape(batch, 160) - cache = torch.randint(0, 256, (batch * 160 + 1, 64, 1, 132), device='cuda', dtype=torch.uint8) - ape = torch.randn(4, 128, device='cuda') - return keys, scores, states, ids, q_lens, kv_lens, cache, blocks, ape, round_scale - - -def candidate_update(keys, scores, states, ids, q_lens, kv_lens, cache, blocks, ape, round_scale): - from lmdeploy.pytorch.backends.cuda.kpool import kpool_prefill_update_cuda - - kpool_prefill_update_cuda(keys, scores, *states, ids, q_lens, kv_lens, cache, blocks, ape, 4, round_scale) - - -@pytest.mark.parametrize('lengths,histories', [ - ([0], [3]), ([0, 0], [0, 3]), ([1, 2], [0, 0]), - ([0, 1, 3, 4, 5], [0, 3, 1, 4, 7]), ([511, 513], [3, 256]), ([8192], [0]), -]) -@pytest.mark.parametrize('round_scale', [False, True]) -@pytest.mark.parametrize('metadata_dtype', [torch.int32, torch.int64]) -def test_kpool_ragged_prefill_matches_reference(lengths, histories, round_scale, metadata_dtype): - args = make_case(lengths, histories, round_scale) - keys, scores, states, ids, q_lens, kv_lens, cache, blocks, ape, _ = args - q_lens, kv_lens = q_lens.to(metadata_dtype), kv_lens.to(metadata_dtype) - expected_states = tuple(state.clone() for state in states) - expected_cache = cache.clone() - reference_update(keys, scores, expected_states, ids, q_lens, kv_lens, expected_cache, blocks, ape, round_scale) - candidate_update(keys, scores, states, ids, q_lens, kv_lens, cache, blocks, ape, round_scale) - torch.testing.assert_close(cache, expected_cache, rtol=0, atol=0) - for actual, expected in zip(states, expected_states): - torch.testing.assert_close(actual, expected, rtol=0, atol=0) - - -def test_kpool_ragged_prefill_graph_changes_layout_and_padding(): - args = make_case([3, 5, 0], [0, 1, 2], True) - keys, scores, states, ids, q_lens, kv_lens, cache, blocks, ape, _ = args - ids[-1] = -1 - initial_states = tuple(state.clone() for state in states) - initial_cache = cache.clone() - candidate_update(*args) - graph = torch.cuda.CUDAGraph() - with torch.cuda.graph(graph): - candidate_update(*args) - for lengths, histories, reordered in [([1, 0, 7], [5, 4, 8], [1, 3, 2]), - ([4, 3, 1], [7, 1, 9], [-1, 2, 1])]: - q_lens.copy_(torch.tensor(lengths, device='cuda')) - kv_lens.copy_(q_lens + torch.tensor(histories, device='cuda')) - ids.copy_(torch.tensor(reordered, device='cuda')) - keys.normal_() - scores.normal_() - expected_states = tuple(state.clone() for state in initial_states) - expected_cache = initial_cache.clone() - reference_update(keys, scores, expected_states, ids, q_lens, kv_lens, expected_cache, blocks, ape, True) - for state, initial in zip(states, initial_states): - state.copy_(initial) - cache.copy_(initial_cache) - graph.replay() - torch.testing.assert_close(cache, expected_cache, rtol=0, atol=0) - for actual, expected in zip(states, expected_states): - torch.testing.assert_close(actual, expected, rtol=0, atol=0) - - -def test_kpool_prefill_layer_views_and_chunk_continuation(): - keys, scores, states, ids, q_lens, kv_lens, cache, blocks, ape, _ = make_case([3, 5], [6, 7], True) - # Runtime states are layer views of [request, layer, pool, width] storage. - banks = tuple(torch.randn(state.size(0), 11, 4, 128, device='cuda', dtype=state.dtype) for state in states) - expected_banks = tuple(bank.clone() for bank in banks) - states = tuple(bank[:, 3] for bank in banks) - expected_states = tuple(bank[:, 3] for bank in expected_banks) - expected_cache = cache.clone() - scores = scores.float() - for lengths in ([3, 5], [1, 7], [6, 2]): - history = kv_lens - q_lens - q_lens.copy_(torch.tensor(lengths, device='cuda')) - kv_lens.copy_(history + q_lens) - keys.normal_() - scores.normal_() - reference_update(keys, scores, expected_states, ids, q_lens, kv_lens, expected_cache, blocks, ape, True) - candidate_update(keys, scores, states, ids, q_lens, kv_lens, cache, blocks, ape, True) - torch.testing.assert_close(cache, expected_cache, rtol=0, atol=0) - for actual, expected in zip(banks, expected_banks): - torch.testing.assert_close(actual, expected, rtol=0, atol=0) - kv_lens.add_(q_lens) - - -def test_indexed_key_scatter_step_major_strides_and_mask(): - from lmdeploy.pytorch.kernels.cuda.fill_kv_cache import fill_indexed_key_cache - from lmdeploy.pytorch.nn.kpool import kpool_packed_cache_views - - torch.manual_seed(59) - cache = torch.randint(0, 256, (17, 64, 1, 132), device='cuda', dtype=torch.uint8) - expected = cache.clone() - keys, scales = kpool_packed_cache_views(cache, 128) - ref_keys, ref_scales = kpool_packed_cache_views(expected, 128) - values = torch.randn(128, 6, device='cuda').to(torch.float8_e4m3fn).T - value_scales = torch.rand(12, device='cuda')[::2] - blocks = torch.arange(1, 17, device='cuda').reshape(2, 8) - groups = torch.tensor([0, -1, 1, 2, -1, 3], device='cuda') - valid = groups >= 0 - for row in [0, 2, 3, 5]: - ref_keys[blocks[row % 2, 0], groups[row]] = values[row] - ref_scales[blocks[row % 2, 0], groups[row], 0] = value_scales[row] - fill_indexed_key_cache(values, value_scales, groups, valid, blocks, keys, scales, page_step=4) - torch.testing.assert_close(cache, expected, rtol=0, atol=0) - - -@pytest.mark.skipif(not torch.cuda.is_available() or torch.cuda.get_device_capability()[0] < 9, - reason='Native FP8 compression requires Hopper or newer.') -@pytest.mark.parametrize('width,pool', [(64, 2), (128, 4)]) -@pytest.mark.parametrize('mode', ['extend', 'decode']) -@pytest.mark.parametrize('round_scale', [False, True]) -def test_kpool_compression_preserves_fp8_bytes_and_scales(width, pool, mode, round_scale): - from lmdeploy.pytorch.kernels.cuda.kpool import compress_kpool - from lmdeploy.pytorch.nn.kpool import kpool_compress, kpool_quantize_fp8 - - torch.manual_seed(73) - keys = torch.randn(257, pool, width, device='cuda', dtype=torch.bfloat16) - scores = torch.randn_like(keys, dtype=torch.float32 if mode == 'decode' else torch.bfloat16) - ape = torch.randn(pool, width, device='cuda') - keys[0].zero_() - keys[1].mul_(1e-6) - keys[2].mul_(1000) - scores[3, 0].fill_(90) - scores[3, 1:].fill_(-90) - expected = kpool_quantize_fp8(kpool_compress(keys, scores, ape, mode=mode), - block_size=width, round_scale=round_scale) - actual = compress_kpool(keys, scores, ape, mode=mode, round_scale=round_scale) - torch.testing.assert_close(actual[0].view(torch.uint8), expected[0].view(torch.uint8), rtol=0, atol=0) - torch.testing.assert_close(actual[1], expected[1], rtol=0, atol=0) diff --git a/tests/pytorch/kernel/test_kpool_selection.py b/tests/pytorch/kernel/test_kpool_selection.py deleted file mode 100644 index 4271b25b12..0000000000 --- a/tests/pytorch/kernel/test_kpool_selection.py +++ /dev/null @@ -1,77 +0,0 @@ -# Copyright (c) OpenMMLab. All rights reserved. -import pytest -import torch - -pytestmark = pytest.mark.skipif(not torch.cuda.is_available() or torch.cuda.get_device_capability()[0] < 9, - reason='DeepGEMM FP8 scoring requires Hopper or newer.') - - -def reference_selection(query, weight, cache, q_lens, kv_lens, blocks): - from lmdeploy.pytorch.backends.cuda.kpool import kpool_score_contiguous_cuda, kpool_select_groups_cuda - from lmdeploy.pytorch.nn.kpool import kpool_expand_selected_groups, kpool_read_packed_cache - - parts = [] - offset = 0 - for request, (q, kv) in enumerate(zip(q_lens.tolist(), kv_lens.tolist())): - seq = kv - q + torch.arange(1, q + 1, device=query.device) - lengths = seq // 4 - keys, scales = kpool_read_packed_cache(cache, blocks[request], kv // 4, 4) - scores = kpool_score_contiguous_cuda(query[offset:offset + q], weight[offset:offset + q], keys, scales, lengths) - selected = kpool_select_groups_cuda(scores, lengths, group_topk=512, max_group_length=kv // 4) - parts.append(kpool_expand_selected_groups(selected, lengths, 4, 2048, seq_lens=seq)) - offset += q - return torch.cat(parts) - - -def make_case(lengths, histories, tied=False): - from lmdeploy.pytorch.nn.kpool import kpool_packed_cache_views - - torch.manual_seed(97) - batch, rows = len(lengths), sum(lengths) - kv = [q + h for q, h in zip(lengths, histories)] - columns = max(1, (max(kv) + 63) // 64) - blocks = (torch.randperm(batch * columns, device='cuda') + 1).reshape(batch, columns) - cache = torch.empty(batch * columns + 1, 64, 1, 132, device='cuda', dtype=torch.uint8) - keys, scales = kpool_packed_cache_views(cache, 128) - keys.copy_(torch.randn(keys.shape, device='cuda').to(torch.float8_e4m3fn)) - scales.uniform_(0.001, 0.05) - query = torch.randn(rows, 32, 128, device='cuda').to(torch.float8_e4m3fn) - weight = torch.zeros(rows, 32, device='cuda') if tied else torch.rand(rows, 32, device='cuda') - return (query, weight, cache, torch.tensor(lengths, device='cuda'), - torch.tensor(kv, device='cuda'), blocks, sum(kv)) - - -@pytest.mark.parametrize('lengths,histories', [ - ([0], [0]), ([1, 2], [0, 0]), ([0, 13, 17], [0, 2079, 6287]), ([511, 513], [3, 4096]), - ([1, 4, 13], [2048, 8189, 32756]), - ([8192], [0]), ([1] * 256, [2] * 256), ([1] * 256, [3] * 256), -]) -@pytest.mark.parametrize('tied', [False, True]) -def test_kpool_prefill_selection_matches_request_loop(lengths, histories, tied): - pytest.importorskip('deep_gemm') - from lmdeploy.pytorch.backends.cuda.kpool import kpool_select_prefill_cuda - - args = make_case(lengths, histories, tied) - expected = reference_selection(*args[:-1]) - actual = kpool_select_prefill_cuda(*args, 4, 2048) - torch.testing.assert_close(actual, expected, rtol=0, atol=0) - - -def test_kpool_prefill_selection_graph_replays_ragged_metadata(): - pytest.importorskip('deep_gemm') - from lmdeploy.pytorch.backends.cuda.kpool import kpool_select_prefill_cuda - - args = make_case([13, 17, 0], [2079, 6278, 0]) - query, weight, cache, q_lens, kv_lens, blocks, capacity = args - kpool_select_prefill_cuda(*args, 4, 2048) - graph = torch.cuda.CUDAGraph() - with torch.cuda.graph(graph): - actual = kpool_select_prefill_cuda(*args, 4, 2048) - for q, kv in [([0, 11, 19], [0, 4097, 4000]), ([17, 0, 13], [3107, 0, 4013])]: - assert sum(kv) <= capacity - q_lens.copy_(torch.tensor(q, device='cuda')) - kv_lens.copy_(torch.tensor(kv, device='cuda')) - weight.uniform_() - expected = reference_selection(query, weight, cache, q_lens, kv_lens, blocks) - graph.replay() - torch.testing.assert_close(actual, expected, rtol=0, atol=0) diff --git a/tests/pytorch/nn/test_glm5_layer_norm.py b/tests/pytorch/nn/test_glm5_layer_norm.py deleted file mode 100644 index 30f9c64f02..0000000000 --- a/tests/pytorch/nn/test_glm5_layer_norm.py +++ /dev/null @@ -1,67 +0,0 @@ -# Copyright (c) OpenMMLab. All rights reserved. -from types import SimpleNamespace - -import pytest -import torch -import torch.nn.functional as F -from torch import nn - -from lmdeploy.pytorch.models import glm5_next -from lmdeploy.pytorch.nn import LayerNorm -from lmdeploy.pytorch.nn.kpool import KPOOL_INDEXER_PARAMETER_NAMES, KPoolIndexer -from lmdeploy.pytorch.weight_loader.model_weight_loader import load_weight - - -@pytest.mark.parametrize('device', [ - 'cpu', pytest.param('cuda', marks=pytest.mark.skipif( - not torch.cuda.is_available(), reason='requires CUDA')), -]) -@pytest.mark.parametrize('dtype', [torch.float16, torch.bfloat16, torch.float32]) -@pytest.mark.parametrize(('owner', 'hidden_size'), [('kpool', 128), ('vision_merger', 128), ('vision_merger', 4096)]) -def test_layer_norm_call_sites_preserve_fp32_contract(monkeypatch, dtype, device, owner, hidden_size): - if owner == 'kpool': - module = KPoolIndexer(128, 2, 128, 8, 4, 4, dtype=dtype, device=device) - assert set(dict(module.named_parameters())) == set(KPOOL_INDEXER_PARAMETER_NAMES) - module.wk = nn.Identity() - norm = module.k_norm - forward = module.project_key - else: - module = glm5_next.Glm5NextVisionPatchMerger( - SimpleNamespace(out_hidden_size=hidden_size, intermediate_size=128, swiglu_limit=10.0), - dtype=dtype, device=device) - # Isolate the real merger's norm -> output cast -> GELU call sequence. - module.proj = nn.Identity() - module.gate_up_proj = nn.Identity() - module.down_proj = nn.Identity() - monkeypatch.setattr(glm5_next, '_glm_swiglu_impl', lambda x, *args, **kwargs: x) - norm = module.post_projection_norm - forward = module - - assert type(norm) is (nn.LayerNorm if owner == 'kpool' else LayerNorm) - assert (norm.eps if owner == 'kpool' else norm.impl.eps) == 1e-6 - assert set(norm.state_dict()) == {'weight', 'bias'} - assert norm.weight.dtype == norm.bias.dtype == torch.float32 - assert not norm.weight.requires_grad and not norm.bias.requires_grad - torch.testing.assert_close(norm.weight, torch.ones_like(norm.weight)) - torch.testing.assert_close(norm.bias, torch.zeros_like(norm.bias)) - generator = torch.Generator().manual_seed(123) - # Non-BF16-representable weights catch accidental parameter downcasting. - load_weight(norm.weight, torch.randn(hidden_size, generator=generator)) - load_weight(norm.bias, torch.randn(hidden_size, generator=generator)) - for tokens in (1, 129, 8192): - inputs = torch.randn(tokens, hidden_size, generator=generator).to(device=device, dtype=dtype) - expected = F.layer_norm(inputs.float(), (hidden_size,), norm.weight, norm.bias, 1e-6).to(dtype) - if owner == 'vision_merger': - expected = F.gelu(expected) - actual = forward(inputs) - assert actual.dtype == dtype - torch.testing.assert_close(actual, expected, rtol=0, atol=0) - - -def test_glm_indexer_uses_default_layer_norm(): - config = SimpleNamespace(hidden_size=16, index_n_heads=2, index_head_dim=8, - index_topk=8, q_lora_rank=4, index_kpool=4) - indexer = glm5_next.Glm5NextSparseAttention._build_indexer( - None, config, layer_idx=0, dtype=torch.bfloat16, device=torch.device('cpu')) - assert type(indexer.k_norm) is nn.LayerNorm - assert indexer.k_norm.weight.dtype == indexer.k_norm.bias.dtype == torch.float32 diff --git a/tests/pytorch/nn/test_moe_options.py b/tests/pytorch/nn/test_moe_options.py deleted file mode 100644 index 49f5d07347..0000000000 --- a/tests/pytorch/nn/test_moe_options.py +++ /dev/null @@ -1,122 +0,0 @@ -# Copyright (c) OpenMMLab. All rights reserved. -import importlib.util -import sys -from types import ModuleType -from unittest.mock import Mock - -import pytest -import torch - - -def test_build_fused_moe_propagates_reduction_options(monkeypatch): - from lmdeploy.pytorch.nn import moe - - captured = {} - - class FakeFusedMoE: - - def __init__(self, **kwargs): - captured.update(kwargs) - - import lmdeploy.pytorch.nn.moe.default as default_moe - monkeypatch.setattr(default_moe, 'FusedMoE', FakeFusedMoE) - - moe.build_fused_moe(16, - 32, - 4, - 2, - quant_config=None, - fp32_acc=True, - output_scale=2.5) - - assert captured['fp32_acc'] is True - assert captured['output_scale'] == 2.5 - - -@pytest.mark.parametrize('options', [ - {'fp32_acc': True}, {'output_scale': 2.5}, - {'fp32_acc': True, 'output_scale': 2.5}, -]) -@pytest.mark.parametrize('quant_method', ['smooth_quant', 'fp8', 'compressed-tensors']) -def test_build_fused_moe_rejects_unsupported_reduction_options(monkeypatch, quant_method, options): - from lmdeploy.pytorch.nn import moe - - class FakeQuantConfig: - quant_dtype = None - activation_scheme = 'static' - weight_block_size = None - bits = 4 - group_size = 128 - - def get_quant_method(self, prefix, module_kind): - return quant_method - - class FakeContext: - quant_config = FakeQuantConfig() - - monkeypatch.setattr(moe, 'get_build_model_context', lambda: FakeContext()) - - with pytest.raises(NotImplementedError, match='fp32_acc or output_scale'): - moe.build_fused_moe(16, - 32, - 4, - 2, - quant_config={}, - **options) - - -@pytest.mark.skipif(not torch.cuda.is_available(), reason='requires CUDA') -def test_moe_reduce_accumulates_in_fp32_before_scaling_and_casting(): - from lmdeploy.pytorch.kernels.cuda.moe.fused_moe import moe_reduce - - hidden = torch.tensor([[[1.25, -2.5], [3.0, 4.0]]], device='cuda', dtype=torch.bfloat16) - weights = torch.tensor([[0.2, 0.7]], device='cuda', dtype=torch.float32) - actual = moe_reduce(hidden, weights, fp32_acc=True, output_scale=2.5) - expected = ((hidden.float() * weights[..., None]).sum(dim=1) * 2.5).to(hidden.dtype) - - torch.testing.assert_close(actual, expected, rtol=0, atol=0) - - -@pytest.fixture -def dlinfer_moe(monkeypatch): - # Exercise the real builder without requiring a DLINFER vendor runtime. - import lmdeploy.pytorch.backends.dlinfer as backend - - kernels = ModuleType('lmdeploy.pytorch.kernels.dlinfer') - for name in ('DlinferMoECommType', 'DlinferMoeMetadata', 'fused_moe', - 'fused_moe_w8a8', 'moe_gating_topk_softmax'): - setattr(kernels, name, Mock()) - spec = importlib.util.spec_from_file_location( - f'{backend.__name__}._test_moe', f'{backend.__path__[0]}/moe.py') - module = importlib.util.module_from_spec(spec) - with monkeypatch.context() as patch: - patch.setitem(sys.modules, kernels.__name__, kernels) - spec.loader.exec_module(module) - return module - - -def _dlinfer_build_spec(**options): - from lmdeploy.pytorch.backends.moe import FusedMoEBuildSpec - - return FusedMoEBuildSpec(top_k=2, num_experts=4, renormalize=True, - hidden_dim=16, ep_size=1, ep_group=None, - layer_idx=0, output_dtype=torch.bfloat16, - num_max_dispatch_tokens_per_rank=32, **options) - - -def test_dlinfer_moe_accepts_default_reduction_options(dlinfer_moe): - impl = dlinfer_moe._build_fused_moe(_dlinfer_build_spec()) - assert isinstance(impl, dlinfer_moe.DlinferFusedMoEImpl) - assert (impl.top_k, impl.num_experts, impl.renormalize, impl.ep_size) == (2, 4, True, 1) - - -@pytest.mark.parametrize('options', [ - {'fp32_acc': True}, {'output_scale': 2.5}, - {'fp32_acc': True, 'output_scale': 2.5}, -]) -def test_dlinfer_moe_rejects_unsupported_reduction_options(dlinfer_moe, monkeypatch, options): - constructor = Mock() - monkeypatch.setattr(dlinfer_moe, 'DlinferFusedMoEImpl', constructor) - with pytest.raises(NotImplementedError, match='fp32_acc or output_scale'): - dlinfer_moe._build_fused_moe(_dlinfer_build_spec(**options)) - constructor.assert_not_called() diff --git a/tests/pytorch/nn/test_rotary_embedding.py b/tests/pytorch/nn/test_rotary_embedding.py index 9ecaabde31..00964fac86 100644 --- a/tests/pytorch/nn/test_rotary_embedding.py +++ b/tests/pytorch/nn/test_rotary_embedding.py @@ -1,4 +1,3 @@ -import pytest import torch from transformers import PretrainedConfig @@ -227,69 +226,3 @@ def test_default_apply_rotary_complex_accepts_half_width_tables_with_empty_key() torch.testing.assert_close(q_embed, _complex_rope_reference(q_states, cos, sin)) assert k_embed.shape == k_states.shape - - -@pytest.mark.parametrize('device', [ - 'cpu', pytest.param('cuda', marks=pytest.mark.skipif( - not torch.cuda.is_available(), reason='requires CUDA')), -]) -@pytest.mark.parametrize('dtype', [torch.float16, torch.bfloat16, torch.float32]) -@pytest.mark.parametrize('enable_fp32_compute', [False, True]) -@pytest.mark.parametrize('inplace', [False, True]) -@pytest.mark.parametrize('complex_mode', [False, True]) -def test_apply_rotary_compute_precision(monkeypatch, dtype, device, enable_fp32_compute, inplace, complex_mode): - from lmdeploy.pytorch.backends.default.op_backend import DefaultOpsBackend - from lmdeploy.pytorch.nn import rotary_embedding - - if device == 'cpu': - monkeypatch.setattr(rotary_embedding, 'get_backend', lambda: DefaultOpsBackend) - module = rotary_embedding.ApplyRotaryEmb(enable_fp32_compute=enable_fp32_compute) - - generator = torch.Generator().manual_seed(123) - # Unequal head counts and strided Q/K exercise the fused CUDA kernel. - query = torch.randn(33, 3, 32, generator=generator).to(device=device, dtype=dtype)[..., ::2] - key = torch.randn(33, 2, 32, generator=generator).to(device=device, dtype=dtype)[..., ::2] - # FP32 tables also catch accidental downcasting before FP32 arithmetic. - cos_value = 0.625 + (2**-12 if enable_fp32_compute else 0) - sin_value = 0.375 + (2**-13 if enable_fp32_compute else 0) - table_dtype = torch.float32 if enable_fp32_compute else dtype - table_dim = 8 if complex_mode else 16 - cos = torch.full((33, table_dim), cos_value, dtype=table_dtype, device=device) - sin = torch.full((33, table_dim), sin_value, dtype=table_dtype, device=device) - original = (query.clone(), key.clone()) - outputs = module(query, key, cos, sin, inplace=inplace, complex_mode=complex_mode) - - for value, saved, actual in zip((query, key), original, outputs): - inputs = saved.float() if enable_fp32_compute else saved - if complex_mode: - rotated = _rotate_complex(inputs) - else: - left, right = inputs.chunk(2, dim=-1) - rotated = torch.cat((-right, left), dim=-1) - expected = (inputs * cos_value + rotated * sin_value).to(dtype) - assert actual.dtype == dtype - torch.testing.assert_close(actual, expected, - rtol=1e-6 if dtype == torch.float32 else 0, - atol=1e-7 if dtype == torch.float32 else 0) - if inplace: - assert actual is value - else: - torch.testing.assert_close(value, saved, rtol=0, atol=0) - - -def test_dlinfer_rejects_fp32_rotary(): - from lmdeploy.pytorch.backends.apply_rotary_emb import ApplyRotaryEmbBuildSpec - from lmdeploy.pytorch.backends.dlinfer.op_backend import DlinferOpsBackend - - with pytest.raises(NotImplementedError, match='enable_fp32_compute=True'): - DlinferOpsBackend.build_op(ApplyRotaryEmbBuildSpec(enable_fp32_compute=True)) - - -def test_glm_vision_uses_common_fp32_rotary(): - from lmdeploy.pytorch.models.glm5_next import Glm5NextVisionAttention - from lmdeploy.pytorch.nn import ApplyRotaryEmb - - config = PretrainedConfig(hidden_size=128, num_heads=2, attention_bias=False) - module = Glm5NextVisionAttention(config, dtype=torch.bfloat16, device='cpu') - assert type(module.apply_rotary_pos_emb) is ApplyRotaryEmb - assert module.apply_rotary_pos_emb.impl.enable_fp32_compute From 18d01883f9cd9ced38e668199883c00cb35a25a6 Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Thu, 24 Sep 2026 03:08:22 +0000 Subject: [PATCH 22/39] refactor: make MoE weighted reduction use FP32 accumulation --- lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py | 5 ----- lmdeploy/pytorch/backends/cuda/moe/default.py | 12 +++--------- lmdeploy/pytorch/backends/dlinfer/moe.py | 4 ++-- lmdeploy/pytorch/backends/moe.py | 2 -- .../kernels/cuda/compressed_tensors_w4a16.py | 4 ++-- lmdeploy/pytorch/kernels/cuda/moe/blocked_fp8.py | 3 +-- lmdeploy/pytorch/kernels/cuda/moe/fused_moe.py | 14 +++----------- lmdeploy/pytorch/models/deepseek_v2.py | 2 -- lmdeploy/pytorch/models/glm5_next.py | 1 - lmdeploy/pytorch/nn/moe/__init__.py | 15 ++++++--------- lmdeploy/pytorch/nn/moe/blocked_fp8.py | 2 -- lmdeploy/pytorch/nn/moe/default.py | 2 -- 12 files changed, 17 insertions(+), 49 deletions(-) diff --git a/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py b/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py index f3bf8ebd2d..006659c6e5 100644 --- a/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py +++ b/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py @@ -290,7 +290,6 @@ def __init__(self, renormalize: bool = False, block_size: int = 128, out_dtype: torch.dtype = torch.float16, - fp32_acc: bool = False, output_scale: float = 1.0): super().__init__() self.num_experts = num_experts @@ -298,7 +297,6 @@ def __init__(self, self.renormalize = renormalize self.block_size = block_size self.out_dtype = out_dtype - self.fp32_acc = fp32_acc self.output_scale = output_scale def ep_expert_list(self, world_size: int, rank: int): @@ -349,7 +347,6 @@ def forward(self, num_experts=num_experts, renormalize=self.renormalize, act_func=act_func, - fp32_acc=self.fp32_acc, output_scale=self.output_scale) output = output.unflatten(0, input_size[:-1]) return output @@ -546,7 +543,6 @@ def fusedmoe_build(self, low_latency_mode: bool = False): def _build_fused_moe_blocked_f8(spec: FusedMoEBlockedF8BuildSpec) -> FusedMoEBlockedF8Impl: """Build a CUDA blocked-FP8 fused MoE implementation.""" if spec.ep_size > 1: - assert not spec.fp32_acc, 'FP32 MoE reduction is not supported by the DeepEP backend yet.' assert spec.output_scale == 1.0, 'MoE output scaling is not supported by the DeepEP backend yet.' assert not spec.custom_gateup_act, 'Custom gate up activation is not supported in EP MoE.' impl = FusedDeepEpMoEBlockedF8Impl( @@ -569,7 +565,6 @@ def _build_fused_moe_blocked_f8(spec: FusedMoEBlockedF8BuildSpec) -> FusedMoEBlo renormalize=spec.renormalize, block_size=spec.block_size, out_dtype=spec.output_dtype, - fp32_acc=spec.fp32_acc, output_scale=spec.output_scale, ) impl.set_scale_fmt(spec.scale_fmt) diff --git a/lmdeploy/pytorch/backends/cuda/moe/default.py b/lmdeploy/pytorch/backends/cuda/moe/default.py index 00af6d41d3..fe60d1c859 100644 --- a/lmdeploy/pytorch/backends/cuda/moe/default.py +++ b/lmdeploy/pytorch/backends/cuda/moe/default.py @@ -26,12 +26,10 @@ def __init__(self, top_k: int, num_experts: int, renormalize: bool = False, - fp32_acc: bool = False, output_scale: float = 1.0): self.num_experts = num_experts self.top_k = top_k self.renormalize = renormalize - self.fp32_acc = fp32_acc self.output_scale = output_scale def update_weights(self, gate_up_weights: torch.Tensor, down_weights: torch.Tensor): @@ -75,7 +73,6 @@ def forward(self, num_experts=num_experts, renormalize=self.renormalize, act_func=act_func, - fp32_acc=self.fp32_acc, output_scale=self.output_scale) @@ -384,15 +381,14 @@ def __init__( num_experts: int, hidden_dim: int, renormalize: bool = False, - fp32_acc: bool = False, output_scale: float = 1.0, layer_idx: int = 0, out_dtype: torch.dtype = torch.bfloat16, num_max_dispatch_tokens_per_rank: int = 128, ): - super().__init__(top_k, num_experts, renormalize, fp32_acc, output_scale) - if fp32_acc or output_scale != 1.0: - raise NotImplementedError('DeepEP MoE does not support fp32_acc or output_scale.') + super().__init__(top_k, num_experts, renormalize, output_scale) + if output_scale != 1.0: + raise NotImplementedError('DeepEP MoE does not support output_scale.') self.num_experts = num_experts self.ep_size = ep_size self.ep_group = ep_group @@ -553,7 +549,6 @@ def _build_fused_moe(spec: FusedMoEBuildSpec) -> FusedMoEImpl: num_experts=spec.num_experts, hidden_dim=spec.hidden_dim, renormalize=spec.renormalize, - fp32_acc=spec.fp32_acc, output_scale=spec.output_scale, layer_idx=spec.layer_idx, out_dtype=spec.output_dtype, @@ -563,6 +558,5 @@ def _build_fused_moe(spec: FusedMoEBuildSpec) -> FusedMoEImpl: top_k=spec.top_k, num_experts=spec.num_experts, renormalize=spec.renormalize, - fp32_acc=spec.fp32_acc, output_scale=spec.output_scale, ) diff --git a/lmdeploy/pytorch/backends/dlinfer/moe.py b/lmdeploy/pytorch/backends/dlinfer/moe.py index 0cbd34a546..c7f4ebbeca 100644 --- a/lmdeploy/pytorch/backends/dlinfer/moe.py +++ b/lmdeploy/pytorch/backends/dlinfer/moe.py @@ -117,9 +117,9 @@ def forward(self, def _build_fused_moe(spec: FusedMoEBuildSpec) -> FusedMoEImpl: """Build a DLINFER fused MoE implementation.""" - if spec.fp32_acc or spec.output_scale != 1.0: + if spec.output_scale != 1.0: raise NotImplementedError( - 'DLINFER fused MoE does not support fp32_acc or output_scale.') + 'DLINFER fused MoE does not support output_scale.') return DlinferFusedMoEImpl( top_k=spec.top_k, num_experts=spec.num_experts, diff --git a/lmdeploy/pytorch/backends/moe.py b/lmdeploy/pytorch/backends/moe.py index 9872aada3d..71ce73812f 100644 --- a/lmdeploy/pytorch/backends/moe.py +++ b/lmdeploy/pytorch/backends/moe.py @@ -73,7 +73,6 @@ class FusedMoEBuildSpec(BuildSpec[FusedMoEImpl]): layer_idx: int output_dtype: torch.dtype num_max_dispatch_tokens_per_rank: int - fp32_acc: bool = False output_scale: float = 1.0 @@ -258,7 +257,6 @@ class FusedMoEBlockedF8BuildSpec(BuildSpec[FusedMoEBlockedF8Impl]): layer_idx: int custom_gateup_act: bool scale_fmt: str | None - fp32_acc: bool = False output_scale: float = 1.0 diff --git a/lmdeploy/pytorch/kernels/cuda/compressed_tensors_w4a16.py b/lmdeploy/pytorch/kernels/cuda/compressed_tensors_w4a16.py index b84c9b25b8..18d1279379 100644 --- a/lmdeploy/pytorch/kernels/cuda/compressed_tensors_w4a16.py +++ b/lmdeploy/pytorch/kernels/cuda/compressed_tensors_w4a16.py @@ -915,7 +915,7 @@ def fused_moe_w4a16( num_bits=num_bits, group_size=group_size, ) - return moe_reduce(expert_output, topk_weights, fp32_acc=True) + return moe_reduce(expert_output, topk_weights) # PyTorch promotes int32 cumsums to int64. Normalize only paths that use # the shared sorter so all routing metadata keeps a homogeneous dtype. @@ -1012,4 +1012,4 @@ def fused_moe_w4a16( ) if valid_routes is not None: expert_output.masked_fill_(~valid_routes[..., None], 0) - return moe_reduce(expert_output, topk_weights, fp32_acc=True) + return moe_reduce(expert_output, topk_weights) diff --git a/lmdeploy/pytorch/kernels/cuda/moe/blocked_fp8.py b/lmdeploy/pytorch/kernels/cuda/moe/blocked_fp8.py index b9dc343fb6..c5fe21a9a8 100644 --- a/lmdeploy/pytorch/kernels/cuda/moe/blocked_fp8.py +++ b/lmdeploy/pytorch/kernels/cuda/moe/blocked_fp8.py @@ -697,7 +697,6 @@ def fused_moe_blocked_fp8(input: torch.Tensor, renormalize: bool = False, act_func: Callable = None, *, - fp32_acc: bool = False, output_scale: float = 1.0) -> torch.Tensor: """Fused moe.""" device = input.device @@ -825,5 +824,5 @@ def fused_moe_blocked_fp8(input: torch.Tensor, ) ret = moe_reduce(intermediate_cache2, topk_weights, - fp32_acc=fp32_acc, output_scale=output_scale) + output_scale=output_scale) return ret diff --git a/lmdeploy/pytorch/kernels/cuda/moe/fused_moe.py b/lmdeploy/pytorch/kernels/cuda/moe/fused_moe.py index ea11c12968..44b548443a 100644 --- a/lmdeploy/pytorch/kernels/cuda/moe/fused_moe.py +++ b/lmdeploy/pytorch/kernels/cuda/moe/fused_moe.py @@ -942,7 +942,6 @@ def _moe_reduce_kernel( stride_wk: tl.constexpr, stride_om, stride_on: tl.constexpr, - fp32_acc: tl.constexpr, K: tl.constexpr, N: tl.constexpr, BLOCK_K: tl.constexpr, @@ -967,11 +966,8 @@ def _moe_reduce_kernel( h = tl.load(h_ptrs, mask=mask_h, other=0.0) w = tl.load(weights_ptrs, mask=mask_k, other=0.0) - if fp32_acc: - h = h.to(tl.float32) - w = w.to(tl.float32) - else: - w = w.to(h.dtype) + h = h.to(tl.float32) + w = w.to(tl.float32) wh = h * w[:, None] o = wh.sum(axis=0) @@ -980,8 +976,7 @@ def _moe_reduce_kernel( tl.store(o_ptrs, o, mask=mask_n) -def moe_reduce(hidden_states: torch.Tensor, topk_weights: torch.Tensor, fp32_acc: bool = False, - *, output_scale: float = 1.0) -> torch.Tensor: +def moe_reduce(hidden_states: torch.Tensor, topk_weights: torch.Tensor, *, output_scale: float = 1.0) -> torch.Tensor: """Weight and reduce experts, optionally scaling before the output cast.""" assert hidden_states.dim() == 3 assert topk_weights.dim() == 2 @@ -1007,7 +1002,6 @@ def moe_reduce(hidden_states: torch.Tensor, topk_weights: torch.Tensor, fp32_acc topk_weights.stride(1), out.stride(0), out.stride(1), - fp32_acc, K, N, BLOCK_K, @@ -1031,7 +1025,6 @@ def fused_moe(hidden_states: torch.Tensor, num_experts: int = None, renormalize: bool = False, act_func: Callable = None, - fp32_acc: bool = False, output_scale: float = 1.0) -> torch.Tensor: """Fused moe.""" M = hidden_states.size(0) @@ -1142,6 +1135,5 @@ def fused_moe(hidden_states: torch.Tensor, ret = moe_reduce(intermediate_cache2, topk_weights, - fp32_acc=fp32_acc, output_scale=output_scale) return ret diff --git a/lmdeploy/pytorch/models/deepseek_v2.py b/lmdeploy/pytorch/models/deepseek_v2.py index f34412a68d..3ebc914112 100644 --- a/lmdeploy/pytorch/models/deepseek_v2.py +++ b/lmdeploy/pytorch/models/deepseek_v2.py @@ -708,7 +708,6 @@ class DeepseekV2MoE(nn.Module): """Deepseek v2 MoE.""" fused_moe_act_func = None - fused_moe_fp32_acc = False fused_moe_output_scale = 1.0 router_routed_scaling_factor = None shared_expert_cls = None @@ -772,7 +771,6 @@ def __init__(self, quant_config=quantization_config, layer_idx=layer_idx, act_func=type(self).fused_moe_act_func, - fp32_acc=type(self).fused_moe_fp32_acc, output_scale=type(self).fused_moe_output_scale, prefix=add_prefix('experts', prefix), ) diff --git a/lmdeploy/pytorch/models/glm5_next.py b/lmdeploy/pytorch/models/glm5_next.py index 2564679bb2..b9da8a639a 100644 --- a/lmdeploy/pytorch/models/glm5_next.py +++ b/lmdeploy/pytorch/models/glm5_next.py @@ -504,7 +504,6 @@ class Glm5NextMoE(DeepseekV2MoE): # Keep only the model-specific clamp; dispatch, quantization and grouped # GEMMs are owned by LMDeploy's generic blocked-FP8 MoE implementation. fused_moe_act_func = staticmethod(_GLM53_COMPACT_FP8_MOE_ACT) - fused_moe_fp32_acc = True # Match the GLM-5.3 contract: routing returns normalized, unscaled weights; # the 2.5 routed scale is applied once to the FP32 expert reduction before # its BF16 store. diff --git a/lmdeploy/pytorch/nn/moe/__init__.py b/lmdeploy/pytorch/nn/moe/__init__.py index 3b30e6a7e2..b38678526f 100644 --- a/lmdeploy/pytorch/nn/moe/__init__.py +++ b/lmdeploy/pytorch/nn/moe/__init__.py @@ -25,7 +25,6 @@ def build_fused_moe( act_func: Callable = None, prefix: str = '', *, - fp32_acc: bool = False, output_scale: float = 1.0, ): """Fused moe builder.""" @@ -48,13 +47,12 @@ def build_fused_moe( all_reduce=all_reduce, layer_idx=layer_idx, act_func=act_func, - fp32_acc=fp32_acc, output_scale=output_scale, ) if quant_method == 'smooth_quant': - if fp32_acc or output_scale != 1.0: - raise NotImplementedError('W8A8 MoE does not support fp32_acc or output_scale.') + if output_scale != 1.0: + raise NotImplementedError('W8A8 MoE does not support output_scale.') assert not bias, 'Quant model does not support bias for now.' assert act_func is None, ('Quant model does not support activation function for now.') from .w8a8 import FusedMoEW8A8 @@ -76,8 +74,8 @@ def build_fused_moe( ) if is_static_per_tensor: - if fp32_acc or output_scale != 1.0: - raise NotImplementedError('Static FP8 MoE does not support fp32_acc or output_scale.') + if output_scale != 1.0: + raise NotImplementedError('Static FP8 MoE does not support output_scale.') assert not bias, ( 'Static FP8 MoE does not support bias.' ) @@ -116,12 +114,11 @@ def build_fused_moe( all_reduce=all_reduce, layer_idx=layer_idx, act_func=act_func, - fp32_acc=fp32_acc, output_scale=output_scale, ) elif quant_method == 'compressed-tensors': - if fp32_acc or output_scale != 1.0: - raise NotImplementedError('W4A16 MoE does not support fp32_acc or output_scale.') + if output_scale != 1.0: + raise NotImplementedError('W4A16 MoE does not support output_scale.') if bias: raise RuntimeError('Compressed-tensors W4A16 routed experts do not support bias.') if act_func is not None: diff --git a/lmdeploy/pytorch/nn/moe/blocked_fp8.py b/lmdeploy/pytorch/nn/moe/blocked_fp8.py index fb9c7fd9f0..11d183e6f8 100644 --- a/lmdeploy/pytorch/nn/moe/blocked_fp8.py +++ b/lmdeploy/pytorch/nn/moe/blocked_fp8.py @@ -158,7 +158,6 @@ def __init__(self, all_reduce: bool = True, layer_idx: int = 0, act_func: Callable = None, - fp32_acc: bool = False, output_scale: float = 1.0): device = device or torch.device('cpu') @@ -192,7 +191,6 @@ def __init__(self, layer_idx=layer_idx, custom_gateup_act=act_func is not None, scale_fmt=scale_fmt, - fp32_acc=fp32_acc, output_scale=output_scale, ), enable_deterministic=build_ctx.enable_deterministic, diff --git a/lmdeploy/pytorch/nn/moe/default.py b/lmdeploy/pytorch/nn/moe/default.py index 4906d967cc..a54d42e233 100644 --- a/lmdeploy/pytorch/nn/moe/default.py +++ b/lmdeploy/pytorch/nn/moe/default.py @@ -131,7 +131,6 @@ def __init__(self, all_reduce: bool = True, layer_idx: int = 0, act_func: Callable = None, - fp32_acc: bool = False, output_scale: float = 1.0): device = device or torch.device('cpu') @@ -160,7 +159,6 @@ def __init__(self, layer_idx=layer_idx, output_dtype=torch.bfloat16, num_max_dispatch_tokens_per_rank=build_ctx.deep_ep_max_tokens_per_rank, - fp32_acc=fp32_acc, output_scale=output_scale, ), enable_deterministic=build_ctx.enable_deterministic, From 1a577dff4ba59f5be324491ec265d42297108223 Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Thu, 24 Sep 2026 03:53:41 +0000 Subject: [PATCH 23/39] refactor: reuse FlashMLA sparse attention for GLM-5.3 --- .../cuda/attention/tilelang_sparse_mla.py | 159 ------------------ lmdeploy/pytorch/models/glm5_next.py | 28 +-- 2 files changed, 15 insertions(+), 172 deletions(-) delete mode 100644 lmdeploy/pytorch/backends/cuda/attention/tilelang_sparse_mla.py diff --git a/lmdeploy/pytorch/backends/cuda/attention/tilelang_sparse_mla.py b/lmdeploy/pytorch/backends/cuda/attention/tilelang_sparse_mla.py deleted file mode 100644 index f373067c58..0000000000 --- a/lmdeploy/pytorch/backends/cuda/attention/tilelang_sparse_mla.py +++ /dev/null @@ -1,159 +0,0 @@ -# Copyright (c) OpenMMLab. All rights reserved. -"""TileLang sparse-MLA backend with reusable logical-index mapping.""" - -from __future__ import annotations - -from typing import Any - -import torch - -from lmdeploy.pytorch.kernels.cuda.sparse_mla_tilelang import ( - sparse_mla_bf16_fwd, -) - -from .sparse_mla import FlashMLAIndexMapper - - -class TilelangSparseMLADecode: - """Run sparse MLA from chronological or model-provided logical indices.""" - - def __init__(self, index_topk: int, index_kpool: int = 1): - if index_topk <= 0: - raise ValueError(f'index_topk must be positive, got {index_topk}') - if index_kpool <= 0: - raise ValueError(f'index_kpool must be positive, got {index_kpool}') - self.index_topk = index_topk - self.index_kpool = index_kpool - tail_width = index_kpool - 1 - self.kernel_topk = ((index_topk + tail_width + 63) // 64) * 64 - self.index_mapper = FlashMLAIndexMapper.build() - - def _pad_logical_indices(self, indices: torch.Tensor) -> torch.Tensor: - """Pad model indices with -1 to TileLang's 64-entry tile width.""" - if indices.ndim != 2: - raise ValueError( - 'Sparse MLA logical indices must have shape [tokens, topk], ' - f'got {tuple(indices.shape)}.') - if indices.size(1) > self.kernel_topk: - raise ValueError( - f'Logical index width {indices.size(1)} exceeds the ' - f'kernel width {self.kernel_topk}.') - if indices.dtype != torch.int32: - indices = indices.to(torch.int32) - return torch.nn.functional.pad( - indices, (0, self.kernel_topk - indices.size(1)), value=-1) - - def _build_physical_indices(self, query: torch.Tensor, - k_cache: torch.Tensor, - attn_metadata: Any) -> torch.Tensor: - kv_seqlens = attn_metadata.kv_seqlens - block_offsets = attn_metadata.block_offsets - batch_size = kv_seqlens.numel() - if query.size(0) != batch_size: - raise NotImplementedError( - 'TileLang sparse MLA currently requires one decode token per ' - f'request, got {query.size(0)} tokens for {batch_size} requests.') - if int(attn_metadata.max_kv_seqlen) > self.index_topk: - raise NotImplementedError( - 'Contexts longer than index_topk require the model KPool indexer.') - - logical = torch.arange(self.kernel_topk, - dtype=torch.int32, - device=query.device) - logical = logical.unsqueeze(0).expand(batch_size, -1).clone() - logical.masked_fill_(logical >= kv_seqlens[:, None], -1) - return self.index_mapper.map_paged_decode( - logical, - block_offsets, - max_q_seqlen=1, - block_size=k_cache.size(1), - ) - - def _map_decode_indices(self, logical_indices: torch.Tensor, - k_cache: torch.Tensor, - attn_metadata: Any) -> torch.Tensor: - logical_indices = self._pad_logical_indices(logical_indices) - batch_size = attn_metadata.kv_seqlens.numel() - if logical_indices.size(0) % batch_size: - raise ValueError( - 'Decode logical-index rows must be divisible by batch size, ' - f'got rows={logical_indices.size(0)}, batch={batch_size}.') - max_q_seqlen = logical_indices.size(0) // batch_size - return self.index_mapper.map_paged_decode( - logical_indices, - attn_metadata.block_offsets, - max_q_seqlen=max_q_seqlen, - block_size=k_cache.size(1), - ).flatten(0, 1)[:, None] - - def _map_prefill_indices(self, logical_indices: torch.Tensor, - attn_metadata: Any) -> torch.Tensor: - logical_indices = self._pad_logical_indices(logical_indices) - return self.index_mapper.map_flat_prefill( - logical_indices, - attn_metadata.q_seqlens, - attn_metadata.cu_seqlens_k, - ) - - def forward(self, - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - k_cache: torch.Tensor, - v_cache: torch.Tensor, - attn_metadata: Any, - scale: float, - cache_writer: Any, - k_scales_zeros: torch.Tensor | None = None, - v_scales_zeros: torch.Tensor | None = None, - logical_indices: torch.Tensor | None = None) -> torch.Tensor: - """Append latent KV, then run BF16 sparse MLA decode.""" - if k_cache.dtype != torch.bfloat16: - raise TypeError('TileLang sparse MLA requires a BF16 KV cache.') - cache_writer._lazy_init(query.device) - cache_impl = cache_writer.impl - max_q_seqlen = cache_impl._get_max_q_seqlen(query, attn_metadata) - cache_impl._fill_kv_cache_impl( - key, - value, - k_cache, - v_cache, - attn_metadata, - max_q_seqlen, - k_scales_zeros=k_scales_zeros, - v_scales_zeros=v_scales_zeros, - ) - if logical_indices is None: - indices = self._build_physical_indices( - query, k_cache, attn_metadata).flatten(0, 1)[:, None] - else: - indices = self._map_decode_indices( - logical_indices, k_cache, attn_metadata) - flat_cache = k_cache.flatten(0, 1) - return sparse_mla_bf16_fwd(query, flat_cache, indices, scale) - - def forward_prefill( - self, - query: torch.Tensor, - key: torch.Tensor, - k_cache: torch.Tensor, - attn_metadata: Any, - scale: float, - cache_writer: Any, - logical_indices: torch.Tensor, - k_scales_zeros: torch.Tensor | None = None, - v_scales_zeros: torch.Tensor | None = None, - ) -> torch.Tensor: - """Append/flatten latent KV, then run sparse MLA prefill.""" - if k_cache.dtype != torch.bfloat16: - raise TypeError('TileLang sparse MLA requires a BF16 KV cache.') - flat_cache = cache_writer.fill_and_flatten_latent_kv_cache( - key, - k_cache, - attn_metadata, - out_dtype=query.dtype, - k_scales_zeros=k_scales_zeros, - v_scales_zeros=v_scales_zeros, - ) - indices = self._map_prefill_indices(logical_indices, attn_metadata) - return sparse_mla_bf16_fwd(query, flat_cache, indices, scale) diff --git a/lmdeploy/pytorch/models/glm5_next.py b/lmdeploy/pytorch/models/glm5_next.py index b9da8a639a..3560e11405 100644 --- a/lmdeploy/pytorch/models/glm5_next.py +++ b/lmdeploy/pytorch/models/glm5_next.py @@ -13,9 +13,7 @@ from torch import distributed as dist from torch import nn -from lmdeploy.pytorch.backends.cuda.attention.tilelang_sparse_mla import ( - TilelangSparseMLADecode, -) +from lmdeploy.pytorch.backends.cuda.attention.sparse_mla import FlashMLASparseImpl from lmdeploy.pytorch.backends.cuda.kpool import ( kpool_compress_quantize_cuda, kpool_prefill_update_cuda, @@ -797,9 +795,15 @@ def __init__(self, v_head_dim=self.v_head_dim, causal=True, ) - self.decode_attn_fwd = TilelangSparseMLADecode( - index_topk=self.index_topk, - index_kpool=getattr(config, 'index_kpool', 1), + dense_mla_impl = self.attn_fwd.impl + self.decode_attn_fwd = FlashMLASparseImpl( + mla_index_topk=self.index_topk, + num_heads=dense_mla_impl.num_heads, + head_size=dense_mla_impl.head_size, + scale=dense_mla_impl.scale, + num_kv_heads=dense_mla_impl.num_kv_heads, + v_head_size=dense_mla_impl.v_head_size, + use_fa3=getattr(dense_mla_impl, 'use_fa3', False), ) def _build_indexer(self, config: Any, layer_idx: int, dtype: torch.dtype, @@ -1267,18 +1271,18 @@ def forward( query_states = self._absorbed_query( unabsorbed_query, num_heads) - attn_output = self.decode_attn_fwd.forward_prefill( + attn_output = self.decode_attn_fwd.forward( query_states, key_states, + key_states[..., :nope_size], past_key_value[0], + past_key_value[0][..., :nope_size], attn_metadata, - scale=self.softmax_scale, - cache_writer=self.attn_fwd, - logical_indices=logical_indices, k_scales_zeros=(None if len(past_key_value) == 2 else past_key_value[2]), v_scales_zeros=(None if len(past_key_value) == 2 else past_key_value[3]), + nsa_indices=logical_indices, ) attn_bmm_out = attn_output.new_empty( q_len, num_heads, self.v_head_dim) @@ -1304,13 +1308,11 @@ def forward( past_key_value[0], past_key_value[0][..., :nope_size], attn_metadata, - scale=self.softmax_scale, - cache_writer=self.attn_fwd, k_scales_zeros=(None if len(past_key_value) == 2 else past_key_value[2]), v_scales_zeros=(None if len(past_key_value) == 2 else past_key_value[3]), - logical_indices=logical_indices, + nsa_indices=logical_indices, ) attn_bmm_out = attn_output.new_empty(q_len, num_heads, self.v_head_dim) self.vc(attn_output, attn_bmm_out) From 8068cd4326bd051ea24d8a210473c0e90b9972a5 Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+users.noreply.github.com> Date: Thu, 24 Sep 2026 04:13:47 +0000 Subject: [PATCH 24/39] revert: remove GLM-specific video sampling --- lmdeploy/vl/media/video.py | 7 +- lmdeploy/vl/media/video_loader.py | 185 +++--------------------------- lmdeploy/vl/model/glm5_next.py | 1 - 3 files changed, 19 insertions(+), 174 deletions(-) diff --git a/lmdeploy/vl/media/video.py b/lmdeploy/vl/media/video.py index c489aa5877..a8c9f2b4ac 100644 --- a/lmdeploy/vl/media/video.py +++ b/lmdeploy/vl/media/video.py @@ -32,17 +32,12 @@ def __init__( self, image_io: ImageMediaIO, num_frames: int = 32, - max_frames: int | None = None, **kwargs, ) -> None: super().__init__() self.image_io = image_io - # ``num_frames`` is LMDeploy's existing request knob; GLM/SGLang call - # the equivalent upper bound ``max_frames``. Accept both at this media - # boundary so model defaults and per-request overrides share one loader. - self.num_frames = int(max_frames if max_frames is not None else - num_frames) + self.num_frames = num_frames # for potential custom arguments from --media-io-kwargs self.kwargs = kwargs diff --git a/lmdeploy/vl/media/video_loader.py b/lmdeploy/vl/media/video_loader.py index 58f98a6727..0cf921f3e5 100644 --- a/lmdeploy/vl/media/video_loader.py +++ b/lmdeploy/vl/media/video_loader.py @@ -18,70 +18,6 @@ logger = get_logger('lmdeploy') -def glm_sample_frame_indices( - total_frames: int, - source_fps: float, - duration: float, - *, - target_fps: float | None = None, - max_frame_count: int | None = None, -) -> list[int]: - """Sample the deterministic temporal pairs expected by GLM video models. - - GLM constructs one visual unit from every two sampled source frames. Its - serving contract therefore differs from the generic uniform sampler in two - ways: the default is 2 FPS with at most 2048 frames, and an odd result - repeats its final frame to keep temporal pairs complete. - """ - if total_frames <= 0: - return [] - target_fps = 2.0 if target_fps is None else float(target_fps) - max_frame_count = (2048 if max_frame_count is None else - int(max_frame_count)) - if target_fps <= 0 or max_frame_count <= 0: - return [] - - max_frame_idx = total_frames - 1 - if not duration: - duration = (round(max_frame_idx / source_fps) + 1 - if source_fps else 0) - extract_t = min(int(duration * target_fps), max_frame_count) - extract_t = max(1, extract_t) - - if source_fps: - duration_per_frame = 1 / source_fps - max_second = int(duration) - indices = [] - current_second = 0.0 - interval = 1 / target_fps - for frame_index in range(total_frames): - timestamp = frame_index * duration_per_frame - if timestamp >= current_second: - current_second += interval - indices.append(frame_index) - if current_second >= max_second: - break - else: - indices = [] - - if len(indices) < extract_t: - start = indices[0] if indices else 0 - end = indices[-1] if indices else max(total_frames - 1, 0) - indices = np.linspace(start, end, extract_t, dtype=int).tolist() - elif len(indices) > extract_t: - indices = np.linspace(0, - total_frames - 1, - extract_t, - dtype=int).tolist() - - # np.linspace can repeat indices for very small inputs. GLM first removes - # those repeats, then pads an odd number of frames with the final sample. - unique_indices = list(dict.fromkeys(int(index) for index in indices)) - if len(unique_indices) & 1: - unique_indices.append(unique_indices[-1]) - return unique_indices - - class VideoLoader: @classmethod @@ -90,27 +26,9 @@ def load_bytes(self, data: bytes, num_frames: int = -1, **kwargs) -> tuple[npt.N raise NotImplementedError @classmethod - def smart_nframes(self, - total_frames_num: int, - num_frames: int, - fps: float, - duration: float, - sampling_strategy: str = 'uniform', - source_fps: float | None = None) -> tuple[int, list[int]]: + def smart_nframes(self, total_frames_num: int, num_frames: int, fps: int, duration: int) -> tuple[int, list[int]]: # resample video to target num_frames and fps # - the minimum of the two will be used - if sampling_strategy == 'glm': - frame_idx = glm_sample_frame_indices( - total_frames_num, - source_fps=source_fps or 0, - duration=duration, - target_fps=None if fps <= 0 else fps, - max_frame_count=None if num_frames <= 0 else num_frames, - ) - return len(frame_idx), frame_idx - if sampling_strategy != 'uniform': - raise ValueError( - f'Unknown video sampling strategy: {sampling_strategy!r}') num_frames_to_sample = total_frames_num if num_frames > 0: num_frames_to_sample = min(num_frames, total_frames_num) @@ -200,28 +118,21 @@ def load_file( self, filepath: Path, num_frames: int = -1, - fps: float = -1, + fps: int = -1, max_duration: int = 300, - sampling_strategy: str = 'uniform', **kwargs, ) -> tuple[npt.NDArray, dict[str, Any]]: with open(filepath, 'rb') as f: data = f.read() - return self.load_bytes(data, - num_frames=num_frames, - fps=fps, - max_duration=max_duration, - sampling_strategy=sampling_strategy, - **kwargs) + return self.load_bytes(data, num_frames=num_frames, fps=fps, max_duration=max_duration, **kwargs) @classmethod def load_bytes( cls, data: bytes, num_frames: int = -1, - fps: float = -1, + fps: int = -1, max_duration: int = 300, - sampling_strategy: str = 'uniform', **kwargs, ) -> tuple[npt.NDArray, dict[str, Any]]: """Load video frames from bytes. @@ -246,35 +157,11 @@ def load_bytes( original_fps = cap.get(cv2.CAP_PROP_FPS) duration = total_frames_num / original_fps if original_fps > 0 else 0 - _, frame_idx = cls.smart_nframes( - total_frames_num, - num_frames, - fps, - duration, - sampling_strategy=sampling_strategy, - source_fps=original_fps, - ) - if not frame_idx: - raise ValueError('Video sampling produced no frame indices.') - - unique_frame_indices = list(dict.fromkeys(frame_idx)) - frame_idx_set = set(unique_frame_indices) - frames, _, valid_frame_indices = cls._read_frames( - cap, - frame_idx_set, - len(unique_frame_indices), - max(frame_idx), - ) - # The GLM sampler may repeat the final frame to complete a temporal - # pair. OpenCV decodes each source index once, so restore the requested - # order (including repeats) after decoding. - decoded = dict(zip(valid_frame_indices, frames)) - ordered_indices = [index for index in frame_idx if index in decoded] - if ordered_indices: - frames = np.stack([decoded[index] for index in ordered_indices]) - else: - frames = frames[:0] - valid_frame_indices = ordered_indices + num_frames_to_sample, frame_idx = cls.smart_nframes(total_frames_num, num_frames, fps, duration) + + frame_idx_set = set(frame_idx) + frames, valid_num_frames, valid_frame_indices = cls._read_frames(cap, frame_idx_set, num_frames_to_sample, + max(frame_idx)) # Use transformers transformers.video_utils.VideoMetadata format # For models like Qwen3-VL/GLM4.5V, this metadata @@ -299,9 +186,8 @@ class DecordVideoLoader(VideoLoader): def load_file(self, filepath: Path, num_frames: int = -1, - fps: float = -1, + fps: int = -1, max_duration: int = 300, - sampling_strategy: str = 'uniform', **kwargs) -> tuple[npt.NDArray, dict[str, Any]]: import decord vr = decord.VideoReader(str(filepath)) @@ -309,16 +195,7 @@ def load_file(self, original_fps = vr.get_avg_fps() duration = total_frames_num / original_fps if original_fps > 0 else 0 - _, frame_idx = self.smart_nframes( - total_frames_num, - num_frames, - fps, - duration, - sampling_strategy=sampling_strategy, - source_fps=original_fps, - ) - if not frame_idx: - raise ValueError('Video sampling produced no frame indices.') + num_frames_to_sample, frame_idx = self.smart_nframes(total_frames_num, num_frames, fps, duration) video = vr.get_batch(frame_idx).asnumpy() # THWC metadata = { @@ -334,9 +211,8 @@ def load_file(self, def load_bytes(self, data: bytes, num_frames: int = -1, - fps: float = -1, + fps: int = -1, max_duration: int = 300, - sampling_strategy: str = 'uniform', **kwargs) -> tuple[npt.NDArray, dict[str, Any]]: tmp_file = tempfile.NamedTemporaryFile(delete=False, suffix='.mp4') try: @@ -346,7 +222,6 @@ def load_bytes(self, num_frames=num_frames, fps=fps, max_duration=max_duration, - sampling_strategy=sampling_strategy, **kwargs) finally: # always cleanup, even if load_file crashes @@ -362,9 +237,8 @@ class TorchCodecVideoLoader(VideoLoader): def load_file(self, filepath: Path, num_frames: int = -1, - fps: float = -1, + fps: int = -1, max_duration: int = 300, - sampling_strategy: str = 'uniform', **kwargs) -> tuple[npt.NDArray, dict[str, Any]]: # torchcodec requires matched ffmpeg, torchcodec, and torch versions # ffmpeg 5.1.2, torch 2.8.0, torchcodec 0.7.0 are verified to work together @@ -376,16 +250,7 @@ def load_file(self, original_fps = decoder.metadata.average_fps duration = total_frames_num / original_fps if original_fps > 0 else 0 - _, frame_idx = self.smart_nframes( - total_frames_num, - num_frames, - fps, - duration, - sampling_strategy=sampling_strategy, - source_fps=original_fps, - ) - if not frame_idx: - raise ValueError('Video sampling produced no frame indices.') + num_frames_to_sample, frame_idx = self.smart_nframes(total_frames_num, num_frames, fps, duration) video = decoder.get_frames_at(frame_idx).data metadata = { @@ -401,9 +266,8 @@ def load_file(self, def load_bytes(self, data: bytes, num_frames: int = -1, - fps: float = -1, + fps: int = -1, max_duration: int = 300, - sampling_strategy: str = 'uniform', **kwargs) -> tuple[npt.NDArray, dict[str, Any]]: tmp_file = tempfile.NamedTemporaryFile(delete=False, suffix='.mp4') try: @@ -413,7 +277,6 @@ def load_bytes(self, num_frames=num_frames, fps=fps, max_duration=max_duration, - sampling_strategy=sampling_strategy, **kwargs) finally: # always cleanup, even if load_file crashes @@ -429,9 +292,8 @@ class TorchVisionVideoLoader(VideoLoader): def load_file(self, filepath: Path, num_frames: int = -1, - fps: float = -1, + fps: int = -1, max_duration: int = 300, - sampling_strategy: str = 'uniform', **kwargs) -> tuple[npt.NDArray, dict[str, Any]]: import torchvision @@ -444,16 +306,7 @@ def load_file(self, original_fps = info['video_fps'] duration = total_frames_num / original_fps if original_fps > 0 else 0 - _, frame_idx = self.smart_nframes( - total_frames_num, - num_frames, - fps, - duration, - sampling_strategy=sampling_strategy, - source_fps=original_fps, - ) - if not frame_idx: - raise ValueError('Video sampling produced no frame indices.') + num_frames_to_sample, frame_idx = self.smart_nframes(total_frames_num, num_frames, fps, duration) video = video[frame_idx] metadata = { @@ -469,9 +322,8 @@ def load_file(self, def load_bytes(self, data: bytes, num_frames: int = -1, - fps: float = -1, + fps: int = -1, max_duration: int = 300, - sampling_strategy: str = 'uniform', **kwargs) -> tuple[npt.NDArray, dict[str, Any]]: tmp_file = tempfile.NamedTemporaryFile(delete=False, suffix='.mp4') try: @@ -481,7 +333,6 @@ def load_bytes(self, num_frames=num_frames, fps=fps, max_duration=max_duration, - sampling_strategy=sampling_strategy, **kwargs) finally: # always cleanup, even if load_file crashes diff --git a/lmdeploy/vl/model/glm5_next.py b/lmdeploy/vl/model/glm5_next.py index 79a3068280..a585c135e7 100644 --- a/lmdeploy/vl/model/glm5_next.py +++ b/lmdeploy/vl/model/glm5_next.py @@ -53,7 +53,6 @@ class GLM5NextVisionModel(VisionModel): 'video': { 'fps': 2.0, 'num_frames': 2048, - 'sampling_strategy': 'glm', }, } From 00f8917c3efd112696eb063a16b0b79169743c83 Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Thu, 24 Sep 2026 07:59:43 +0000 Subject: [PATCH 25/39] revert: remove model media I/O defaults merging --- lmdeploy/serve/processors/multimodal.py | 28 ------------------------- lmdeploy/vl/model/glm5_next.py | 8 ------- 2 files changed, 36 deletions(-) diff --git a/lmdeploy/serve/processors/multimodal.py b/lmdeploy/serve/processors/multimodal.py index 8fc698e28c..19a1bfea2d 100644 --- a/lmdeploy/serve/processors/multimodal.py +++ b/lmdeploy/serve/processors/multimodal.py @@ -54,33 +54,6 @@ def __init__(self, self.backend = backend self.allowed_media_domains = allowed_media_domains - @staticmethod - def merge_media_io_kwargs( - defaults: dict[str, Any] | None, - overrides: dict[str, Any] | None, - ) -> dict[str, Any]: - """Merge model media defaults with request-level modality overrides.""" - merged: dict[str, Any] = {} - for key, value in (defaults or {}).items(): - merged[key] = dict(value) if isinstance(value, dict) else value - for key, value in (overrides or {}).items(): - if isinstance(value, dict) and isinstance(merged.get(key), dict): - merged[key] = {**merged[key], **value} - else: - merged[key] = dict(value) if isinstance(value, dict) else value - return merged - - def resolve_media_io_kwargs( - self, - overrides: dict[str, Any] | None, - ) -> dict[str, Any]: - """Resolve model-specific media defaults without mutating requests.""" - model = getattr(self.vl_encoder, 'model', None) - defaults = getattr(model, 'default_media_io_kwargs', None) - if callable(defaults): - defaults = defaults() - return self.merge_media_io_kwargs(defaults, overrides) - @staticmethod def merge_message_content(msg: dict) -> dict: """Merge multimodal content blocks and ensure content field exists. @@ -442,7 +415,6 @@ async def _get_multimodal_prompt_input(self, """Process multimodal prompt and return processed data for inference engines.""" chat_template = self.chat_template if do_preprocess else BaseChatTemplate() - media_io_kwargs = self.resolve_media_io_kwargs(media_io_kwargs) messages = await self.async_parse_multimodal_item(messages, media_io_kwargs, allowed_media_domains=self.allowed_media_domains) diff --git a/lmdeploy/vl/model/glm5_next.py b/lmdeploy/vl/model/glm5_next.py index a585c135e7..7effc460c9 100644 --- a/lmdeploy/vl/model/glm5_next.py +++ b/lmdeploy/vl/model/glm5_next.py @@ -47,14 +47,6 @@ class GLM5NextVisionModel(VisionModel): """Prepare GLM-5.3 images/videos for the native PyTorch vision tower.""" _arch = ['Glm5NextForConditionalGeneration'] - # Match GLM/SGLang's video contract before the HF processor: sample at - # 2 FPS, cap at 2048 source frames, and complete temporal pairs. - default_media_io_kwargs = { - 'video': { - 'fps': 2.0, - 'num_frames': 2048, - }, - } @classmethod def match(cls, config): From 3f21b23f8e5a94224678f8434cbcf9bfd83f8ccb Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Thu, 24 Sep 2026 13:38:57 +0000 Subject: [PATCH 26/39] fix(pytorch): enable GLM DP and expert parallelism Propagate GLM activation and routed scaling through DeepEP prefill and decode. Use FP32 local expert reduction for non-default output scaling while preserving the default BF16 path and transport. Avoid reducing combined experts twice and keep shared experts unscaled. Pad sparse FlashMLA indices with invalid entries for its tile alignment and accept the current DeepGEMM masked-GEMM symbol. --- .../backends/cuda/attention/sparse_mla.py | 6 ++ .../pytorch/backends/cuda/moe/blocked_fp8.py | 55 ++++++++++++++----- lmdeploy/pytorch/kernels/cuda/moe/ep_fp8.py | 25 ++++++--- lmdeploy/pytorch/models/glm5_next.py | 17 ++++-- lmdeploy/pytorch/nn/moe/blocked_fp8.py | 12 +++- .../pytorch/third_party/deep_gemm/__init__.py | 4 +- 6 files changed, 91 insertions(+), 28 deletions(-) diff --git a/lmdeploy/pytorch/backends/cuda/attention/sparse_mla.py b/lmdeploy/pytorch/backends/cuda/attention/sparse_mla.py index d0d7f56d9b..7f488d2770 100644 --- a/lmdeploy/pytorch/backends/cuda/attention/sparse_mla.py +++ b/lmdeploy/pytorch/backends/cuda/attention/sparse_mla.py @@ -140,6 +140,12 @@ def _flash_mla_sparse_forward(self, query: torch.Tensor, indexed_kv: torch.Tenso if pad_heads: query = torch.nn.functional.pad(query, (0, 0, 0, pad_heads)) + # Sparse FlashMLA consumes pairs of 64-entry index tiles. Invalid + # entries preserve the selected keys when top-k is not tile-aligned. + pad_indices = -indices.size(-1) % 128 + if pad_indices: + indices = torch.nn.functional.pad(indices, (0, pad_indices), value=-1) + attn_output = flash_mla_sparse_fwd(query, indexed_kv, indices, sm_scale=self.scale)[0] return attn_output[:, :num_q_heads] diff --git a/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py b/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py index 006659c6e5..bba12e1821 100644 --- a/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py +++ b/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py @@ -46,6 +46,7 @@ def __init__( num_max_dispatch_tokens_per_rank: int = 128, chunk_size: int | None = 32 * 1024, expert_alignment: int = 128, + fp32_acc: bool = False, ): self.layer_index = layer_index self.top_k = top_k @@ -55,6 +56,7 @@ def __init__( self.out_dtype = out_dtype self.fp8_dtype = fp8_dtype self.scale_fmt = scale_fmt + self.fp32_acc = fp32_acc self.token_dispatcher = DeepEPTokenDispatcherNormal( group=ep_group, num_experts=num_experts, @@ -75,6 +77,7 @@ def forward( down_weights: torch.Tensor, down_scale: torch.Tensor, expert_list: list[int] = None, + act_func: Callable = None, ): hs_quant, hs_scale = per_token_group_quant_fp8(hidden_states, self.block_size, @@ -87,7 +90,8 @@ def forward( expert_list, ) out_states = fused_moe_v3_fp8(x, recv_topk_ids, recv_topk_weights, (up_weights, up_scale), - (down_weights, down_scale), recv_tokens_per_expert) + (down_weights, down_scale), recv_tokens_per_expert, + act_func=act_func, scale_fmt=self.scale_fmt, fp32_acc=self.fp32_acc) return self.token_dispatcher.combine(out_states) def capture(self): @@ -115,9 +119,10 @@ def combine_async(self, x: torch.Tensor, handle: tuple, previous_event=None, asy def release(self): return self.token_dispatcher.release() - def fusedmoe_forward(self, state, up_weight, up_scale, down_weight, down_scale): + def fusedmoe_forward(self, state, up_weight, up_scale, down_weight, down_scale, act_func=None): return fused_moe_v3_fp8(state['recv_hidden_states'], state['recv_topk_idx'], state['recv_topk_weights'], - (up_weight, up_scale), (down_weight, down_scale), state['recv_tokens_per_expert']) + (up_weight, up_scale), (down_weight, down_scale), state['recv_tokens_per_expert'], + act_func=act_func, scale_fmt=self.scale_fmt, fp32_acc=self.fp32_acc) def per_token_group_quant_fp8(self, x: torch.Tensor, @@ -140,11 +145,13 @@ def __init__( block_size: int = 128, out_dtype: torch.dtype = torch.bfloat16, num_max_dispatch_tokens_per_rank: int = 128, + scale_fmt: str | None = None, ): self.num_experts = num_experts self.layer_index = layer_index self.block_size = block_size self.out_dtype = out_dtype + self.scale_fmt = scale_fmt self.token_dispatcher = DeepEPTokenDispatcherLowLatency( group=ep_group, num_experts=num_experts, @@ -168,6 +175,7 @@ def experts( gate_down_scale: torch.Tensor, masked_m: torch.Tensor, expected_m: int, + act_func: Callable = None, ): gate_up_weight_fp8 = (gate_up_weight, gate_up_scale) gate_down_weight_fp8 = (gate_down_weight, gate_down_scale) @@ -186,7 +194,15 @@ def experts( (gateup_output.shape[0], gateup_output.shape[1], gateup_output.shape[2] // 2 // self.block_size), device=gateup_output.device, dtype=torch.float32) - silu_and_mul_masked_post_quant_fwd(gateup_output, down_input, down_input_scale, self.block_size, masked_m) + if act_func is None: + silu_and_mul_masked_post_quant_fwd(gateup_output, down_input, down_input_scale, self.block_size, masked_m) + else: + # Only masked_m valid rows are consumed by the following GEMM. + # Reuse the model's activation and the shared quantizer unchanged. + from lmdeploy.pytorch.kernels.cuda.blocked_gemm_fp8 import _quant_fp8_launcher + activated = act_func(gateup_output.flatten(0, 1)) + _quant_fp8_launcher(activated, self.block_size, down_input.flatten(0, 1), + down_input_scale.flatten(0, 1), scale_fmt=self.scale_fmt) del gateup_output down_output = torch.empty((num_groups, m, gate_down_weight.size(1)), device=down_input.device, @@ -205,11 +221,12 @@ def forward( down_weights: torch.Tensor, down_scale: torch.Tensor, expert_list: list[int] = None, + act_func: Callable = None, ): recv_hidden_states, topk_idx, topk_weights, masked_m, expected_m = self.token_dispatcher.dispatch( hidden_states, topk_ids, topk_weights, self.num_experts) out_states = self.experts(recv_hidden_states, up_weights, up_scale, down_weights, down_scale, masked_m, - expected_m) + expected_m, act_func=act_func) return self.token_dispatcher.combine(out_states, topk_idx, topk_weights) def wait(self, event): @@ -231,14 +248,15 @@ def combine_async(self, async_finish: bool): return self.token_dispatcher.combine_async(hidden_states, topk_idx, topk_weights, handle, async_finish) - def fusedmoe_forward(self, state, up_weight, up_scale, down_weight, down_scale): + def fusedmoe_forward(self, state, up_weight, up_scale, down_weight, down_scale, act_func=None): recv_hidden_states = state['recv_hidden_states'] masked_m = state['recv_expert_count'] hidden_shape = state['raw_hidden_shape'] topk_idx = state['topk_idx'] expected_m = (hidden_shape[0] * self.token_dispatcher.buffer_low_latency.group_size * topk_idx.shape[1] + self.token_dispatcher.num_experts) // self.token_dispatcher.num_experts - return self.experts(recv_hidden_states, up_weight, up_scale, down_weight, down_scale, masked_m, expected_m) + return self.experts(recv_hidden_states, up_weight, up_scale, down_weight, down_scale, masked_m, expected_m, + act_func=act_func) def _build_deepep_moe( @@ -256,6 +274,7 @@ def _build_deepep_moe( num_max_dispatch_tokens_per_rank: int = 128, chunk_size: int | None = 32 * 1024, expert_alignment: int = 128, + fp32_acc: bool = False, ): if low_latency_mode: return FusedMoELowLatency(ep_size=ep_size, @@ -265,6 +284,7 @@ def _build_deepep_moe( layer_index=layer_idx, block_size=block_size, out_dtype=out_dtype, + scale_fmt=scale_fmt, num_max_dispatch_tokens_per_rank=num_max_dispatch_tokens_per_rank) return FusedMoENormal(ep_size=ep_size, ep_group=ep_group, @@ -278,7 +298,8 @@ def _build_deepep_moe( scale_fmt=scale_fmt, num_max_dispatch_tokens_per_rank=num_max_dispatch_tokens_per_rank, chunk_size=chunk_size, - expert_alignment=expert_alignment) + expert_alignment=expert_alignment, + fp32_acc=fp32_acc) class TritonFusedMoEBlockedF8Impl(FusedMoEBlockedF8Impl): @@ -354,6 +375,8 @@ def forward(self, class FusedDeepEpMoEBlockedF8Impl(TritonFusedMoEBlockedF8Impl): + output_scale = 1.0 + def __init__(self, ep_size: int, ep_group: dist.ProcessGroup, @@ -365,8 +388,9 @@ def __init__(self, out_dtype: torch.dtype = torch.bfloat16, fp8_dtype: torch.dtype = torch.float8_e4m3fn, num_max_dispatch_tokens_per_rank: int = 128, - layer_idx: int = 0): - super().__init__(top_k, num_experts, renormalize, block_size, out_dtype) + layer_idx: int = 0, + output_scale: float = 1.0): + super().__init__(top_k, num_experts, renormalize, block_size, out_dtype, output_scale) self.num_experts = num_experts self.ep_size = ep_size self.ep_group = ep_group @@ -431,7 +455,7 @@ def _forward_moe(self, low_latency_mode = step_ctx.global_is_decoding() and self.use_deep_gemm moe = self.fusedmoe_build(low_latency_mode) out_states = moe.forward(hidden_states, topk_weights, topk_ids, gate_up_weights, gate_up_scale, down_weights, - down_scale, expert_list) + down_scale, expert_list, act_func=act_func) out_states = gather_outputs_by_attn_tp(out_states, split_size) return out_states @@ -521,7 +545,10 @@ def piecewise_forward( self._piecewise_forward = piecewise_forward def do_renormalize(self, topk_weights): - return _renormalize(topk_weights, self.renormalize) + weights = _renormalize(topk_weights, self.renormalize) + # DeepEP owns the final distributed combine. Fold only the routed + # scale into FP32 routing weights; the separate shared expert is unscaled. + return weights if self.output_scale == 1.0 else weights.float() * self.output_scale def fusedmoe_build(self, low_latency_mode: bool = False): deepep_moe = _build_deepep_moe(low_latency_mode, @@ -536,6 +563,7 @@ def fusedmoe_build(self, low_latency_mode: bool = False): scale_fmt=self.scale_fmt, layer_idx=self.layer_idx, num_max_dispatch_tokens_per_rank=self.num_max_dispatch_tokens_per_rank, + fp32_acc=self.output_scale != 1.0, chunk_size=16 * 1024) return deepep_moe @@ -543,8 +571,6 @@ def fusedmoe_build(self, low_latency_mode: bool = False): def _build_fused_moe_blocked_f8(spec: FusedMoEBlockedF8BuildSpec) -> FusedMoEBlockedF8Impl: """Build a CUDA blocked-FP8 fused MoE implementation.""" if spec.ep_size > 1: - assert spec.output_scale == 1.0, 'MoE output scaling is not supported by the DeepEP backend yet.' - assert not spec.custom_gateup_act, 'Custom gate up activation is not supported in EP MoE.' impl = FusedDeepEpMoEBlockedF8Impl( ep_size=spec.ep_size, ep_group=spec.ep_group, @@ -557,6 +583,7 @@ def _build_fused_moe_blocked_f8(spec: FusedMoEBlockedF8BuildSpec) -> FusedMoEBlo fp8_dtype=spec.fp8_dtype, num_max_dispatch_tokens_per_rank=spec.num_max_dispatch_tokens_per_rank, layer_idx=spec.layer_idx, + output_scale=spec.output_scale, ) else: impl = TritonFusedMoEBlockedF8Impl( diff --git a/lmdeploy/pytorch/kernels/cuda/moe/ep_fp8.py b/lmdeploy/pytorch/kernels/cuda/moe/ep_fp8.py index 8eab01008f..ab2f39c0b5 100644 --- a/lmdeploy/pytorch/kernels/cuda/moe/ep_fp8.py +++ b/lmdeploy/pytorch/kernels/cuda/moe/ep_fp8.py @@ -1,5 +1,7 @@ # Copyright (c) OpenMMLab. All rights reserved. # modify from dlblas: https://github.com/DeepLink-org/DLBlas +from collections.abc import Callable + import torch import triton import triton.language as tl @@ -152,18 +154,23 @@ def fused_moe_v3_fp8( w13_weight_fp8: tuple[torch.Tensor, torch.Tensor], w2_weight_fp8: tuple[torch.Tensor, torch.Tensor], num_recv_tokens_per_expert: list[int] | None, + act_func: Callable | None = None, + scale_fmt: str | None = None, + fp32_acc: bool = False, ): hidden_states_fp8, hidden_states_scale = hidden_states_fp8 if num_recv_tokens_per_expert is None: - return hidden_states_fp8.to(torch.bfloat16) + return torch.zeros_like(hidden_states_fp8, dtype=torch.bfloat16) all_tokens = sum(num_recv_tokens_per_expert) if all_tokens <= 0: - return hidden_states_fp8.to(torch.bfloat16) + return torch.zeros_like(hidden_states_fp8, dtype=torch.bfloat16) from lmdeploy.pytorch.third_party.deep_gemm import get_mn_major_tma_aligned_tensor m, k = hidden_states_fp8.size() n = w13_weight_fp8[0].size(1) block_size = k // hidden_states_scale.size(1) - gather_out = torch.empty_like(hidden_states_fp8, device=hidden_states_fp8.device, dtype=torch.bfloat16) + # ep_gather already accumulates in its output dtype. Reuse that contract + # instead of introducing a second reduction kernel for FP32 accumulation. + gather_out = torch.empty_like(hidden_states_fp8, dtype=torch.float32 if fp32_acc else torch.bfloat16) input_tensor = torch.empty((all_tokens, k), device=hidden_states_fp8.device, dtype=hidden_states_fp8.dtype) input_tensor_scale = torch.empty((all_tokens, k // block_size), device=hidden_states_fp8.device, @@ -183,12 +190,16 @@ def fused_moe_v3_fp8( input_tensor_scale = get_mn_major_tma_aligned_tensor(input_tensor_scale) _deepgemm_grouped_fp8_nt_contiguous((input_tensor, input_tensor_scale), w13_weight_fp8, gateup_output, m_indices) - down_input = torch.empty((all_tokens, n // 2), device=gateup_output.device, dtype=torch.bfloat16) - silu_and_mul(gateup_output.view(-1, n), down_input) + if act_func is None: + down_input = silu_and_mul(gateup_output.view(-1, n)) + else: + down_input = act_func(gateup_output.view(-1, n)) del gateup_output - down_input_fp8, down_input_scale = per_token_group_quant_fp8(down_input, block_size) + down_input_fp8, down_input_scale = per_token_group_quant_fp8(down_input, block_size, scale_fmt=scale_fmt) down_input_scale = get_mn_major_tma_aligned_tensor(down_input_scale) down_output = torch.empty((all_tokens, k), device=gather_out.device, dtype=torch.bfloat16) _deepgemm_grouped_fp8_nt_contiguous((down_input_fp8, down_input_scale), w2_weight_fp8, down_output, m_indices) ep_gather(down_output, topk_idx, topk_weights, output_index, gather_out) - return gather_out + # DeepEP transports BF16 partial sums. Local expert reduction can be FP32, + # but this is not an all-FP32 cross-rank reduction contract. + return gather_out.to(torch.bfloat16) diff --git a/lmdeploy/pytorch/models/glm5_next.py b/lmdeploy/pytorch/models/glm5_next.py index 3560e11405..ae9d2b6580 100644 --- a/lmdeploy/pytorch/models/glm5_next.py +++ b/lmdeploy/pytorch/models/glm5_next.py @@ -451,8 +451,8 @@ def __init__(self, config: Any): self.scoring_func = config.scoring_func self.renormalize = bool(config.norm_topk_prob and self.top_k > 1) # Keep the generic router output normalized. GLM's model-level 2.5 - # factor is owned by ``fused_moe_output_scale`` below so it is applied - # once, after the FP32 routed-expert reduction. + # factor is owned by ``fused_moe_output_scale`` below. EP1 applies it + # after reduction; DeepEP folds it into FP32 combine weights. self.routed_scaling_factor = 1.0 self.router_n_groups = getattr(config, 'router_n_groups', -1) contract = ( @@ -503,8 +503,8 @@ class Glm5NextMoE(DeepseekV2MoE): # GEMMs are owned by LMDeploy's generic blocked-FP8 MoE implementation. fused_moe_act_func = staticmethod(_GLM53_COMPACT_FP8_MOE_ACT) # Match the GLM-5.3 contract: routing returns normalized, unscaled weights; - # the 2.5 routed scale is applied once to the FP32 expert reduction before - # its BF16 store. + # EP1 scales the FP32 expert reduction before its BF16 store. DeepEP + # scales FP32 combine weights; shared experts remain unscaled. router_routed_scaling_factor = 1.0 fused_moe_output_scale = 2.5 shared_expert_cls = Glm5NextMLP @@ -512,6 +512,15 @@ class Glm5NextMoE(DeepseekV2MoE): def __init__(self, config: Any, layer_idx: int, *args, **kwargs): kwargs.setdefault('prefix', f'model.layers.{layer_idx}.mlp') super().__init__(config, layer_idx, *args, **kwargs) + dist_config = get_dist_manager().current_config() + if dist_config.ep > 1: + # DeepEP already combines routed experts. Reduce only the TP + # shared projection, otherwise the outer TP sum multiplies the + # complete routed contribution by the attention TP size. + self._all_reduce = False + if dist_config.dp == 1 and self.shared_experts is not None: + down_proj = self.shared_experts.down_proj + down_proj.all_reduce = down_proj.tp > 1 # Keep the shared+routed local sum and the generic expert kernels. # Promote only the final TP collective: BF16 collective reduction # order depends on message size (AR versus multi-token verification). diff --git a/lmdeploy/pytorch/nn/moe/blocked_fp8.py b/lmdeploy/pytorch/nn/moe/blocked_fp8.py index 11d183e6f8..157d9273a4 100644 --- a/lmdeploy/pytorch/nn/moe/blocked_fp8.py +++ b/lmdeploy/pytorch/nn/moe/blocked_fp8.py @@ -144,6 +144,8 @@ def weight_loader_with_quant(self, param: torch.nn.Parameter, loaded_weight: tor class FusedMoEBlockedF8(FusedMoEBase): """Fused moe blocked f8.""" + output_scale = 1.0 + def __init__(self, hidden_dim: int, ffn_dim: int, @@ -232,6 +234,7 @@ def __init__(self, self.dtype = dtype self.device = device self.act_func = act_func + self.output_scale = output_scale @staticmethod def _update_args(hidden_dim: int, ffn_dim: int, align: int): @@ -254,6 +257,9 @@ def before_dispatch(self, state: DispatchInputs): state = state.to_dict() moe_type = state['moe_type'] + if moe_type in (MoeType.DSAsyncPrefill, MoeType.DSAsyncDecode) and self.output_scale != 1.0: + # Async callers already normalized before splitting microbatches. + state['topk_weights'] = state['topk_weights'].float() * self.output_scale if moe_type == MoeType.DSAsyncPrefill: fusedmoe = self.fusedmoe_build(low_latency_mode=False) state['fusedmoe'] = fusedmoe @@ -336,7 +342,8 @@ def gemm(self, state: dict): state['recv_hidden_states'] = state['fusedmoe'].fusedmoe_forward(state, self.gate_up.weight, self.gate_up.weight_scale_inv, self.down.weight, - self.down.weight_scale_inv) + self.down.weight_scale_inv, + act_func=self.act_func) gemm_state = { 'fusedmoe': state['fusedmoe'], 'hidden_states': state['recv_hidden_states'], @@ -347,7 +354,8 @@ def gemm(self, state: dict): state['recv_hidden_states'] = state['fusedmoe'].fusedmoe_forward(state, self.gate_up.weight, self.gate_up.weight_scale_inv, self.down.weight, - self.down.weight_scale_inv) + self.down.weight_scale_inv, + act_func=self.act_func) gemm_state = { 'fusedmoe': state['fusedmoe'], 'hidden_states': state['recv_hidden_states'], diff --git a/lmdeploy/pytorch/third_party/deep_gemm/__init__.py b/lmdeploy/pytorch/third_party/deep_gemm/__init__.py index cc9cb4333c..70f7e937af 100644 --- a/lmdeploy/pytorch/third_party/deep_gemm/__init__.py +++ b/lmdeploy/pytorch/third_party/deep_gemm/__init__.py @@ -76,7 +76,9 @@ def m_grouped_fp8_gemm_nt_contiguous(a, b, d, m_indices, recipe=None, compiled_d try: - from deep_gemm import m_grouped_fp8_gemm_nt_masked + m_grouped_fp8_gemm_nt_masked = getattr(deep_gemm, 'm_grouped_fp8_gemm_nt_masked', None) + if m_grouped_fp8_gemm_nt_masked is None: + m_grouped_fp8_gemm_nt_masked = deep_gemm.fp8_m_grouped_gemm_nt_masked except Exception: from deep_gemm import m_grouped_gemm_fp8_fp8_bf16_nt_masked From 38779f12d36299d86fd7918fc62b460e9de90cc1 Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Mon, 28 Sep 2026 04:05:43 +0000 Subject: [PATCH 27/39] perf(pytorch): fuse GLM KPool verification and reuse shared operators --- .../cuda/comm/flashinfer_allreduce.py | 33 +++-- lmdeploy/pytorch/backends/cuda/hc_prepost.py | 9 +- lmdeploy/pytorch/backends/cuda/kda.py | 2 +- lmdeploy/pytorch/backends/cuda/kpool.py | 37 +++++- lmdeploy/pytorch/backends/hc_prepost.py | 4 +- .../pytorch/kernels/cuda/causal_conv1d.py | 10 +- .../pytorch/kernels/cuda/dsv4/hc_prepost.py | 52 ++++++++ .../pytorch/kernels/cuda/gated_delta_rule.py | 4 +- lmdeploy/pytorch/kernels/cuda/kpool.py | 125 +++++++++++++++++- lmdeploy/pytorch/models/glm5_next.py | 74 +++++++---- lmdeploy/pytorch/nn/hc_prepost.py | 4 +- lmdeploy/pytorch/nn/linear/base.py | 4 +- 12 files changed, 294 insertions(+), 64 deletions(-) diff --git a/lmdeploy/pytorch/backends/cuda/comm/flashinfer_allreduce.py b/lmdeploy/pytorch/backends/cuda/comm/flashinfer_allreduce.py index 9cd5a431b5..f4edeb1542 100644 --- a/lmdeploy/pytorch/backends/cuda/comm/flashinfer_allreduce.py +++ b/lmdeploy/pytorch/backends/cuda/comm/flashinfer_allreduce.py @@ -56,9 +56,8 @@ def __init__(self, group: dist.ProcessGroup): """Initialize lazy FlashInfer state for ``group``.""" self.group = group self._comm = None - self._workspace = None - self._hidden_dim = None - self._dtype = None + self._workspaces = {} + self._comm_backend = None self._disabled = False self._world_size = dist.get_world_size(group) major, minor = torch.cuda.get_device_capability() @@ -82,6 +81,7 @@ def _initialize(self) -> bool: try: import flashinfer.comm as comm from flashinfer.comm.cuda_ipc import cudart + from flashinfer.comm.mnnvl import TorchDistBackend if not (hasattr(comm, 'allreduce_fusion') and hasattr( comm, 'create_allreduce_fusion_workspace')): raise ImportError(message) @@ -89,6 +89,7 @@ def _initialize(self) -> bool: # Resolve FlashInfer's CUDA runtime before TileLang loads its # libcudart shim, which the lazy loader could otherwise select. cudart.cudaSetDevice(torch.cuda.current_device()) + self._comm_backend = TorchDistBackend(group=self.group) except Exception as e: self._disable(e) return False @@ -101,11 +102,11 @@ def is_available(self) -> bool: def supports(self, dtype: torch.dtype) -> bool: """Whether this group can handle the given dtype.""" - return dtype in (torch.float16, torch.bfloat16) and self.is_available() + return dtype in (torch.float16, torch.bfloat16, torch.float32) and self.is_available() def _supports_input(self, input: torch.Tensor) -> bool: """Whether a flattened input satisfies the FlashInfer launch limits.""" - return (input.nbytes <= self._max_size + return (input.numel() > 0 and input.nbytes <= self._max_size and input.size(0) <= _MAX_TOKEN_NUM and input.is_contiguous() and self.supports(input.dtype)) @@ -113,9 +114,9 @@ def _supports_input(self, input: torch.Tensor) -> bool: def _get_workspace(self, input: torch.Tensor): """Return the workspace for the input shape and dtype.""" hidden_dim = input.size(-1) - if self._workspace is not None: - assert self._hidden_dim == hidden_dim and self._dtype == input.dtype - return self._workspace + key = (hidden_dim, input.dtype) + if key in self._workspaces: + return self._workspaces[key] rank = dist.get_rank(self.group) max_token_num = min(_MAX_TOKEN_NUM, self._max_size // input[0].nbytes) @@ -127,7 +128,7 @@ def _get_workspace(self, input: torch.Tensor): max_token_num=max_token_num, hidden_dim=hidden_dim, dtype=input.dtype, - group=self.group, + comm_backend=self._comm_backend, ) if workspace is None: raise RuntimeError('workspace creation returned None') @@ -135,13 +136,11 @@ def _get_workspace(self, input: torch.Tensor): self._disable(f'workspace initialization failed: {e}') return None - self._workspace = workspace - self._hidden_dim = hidden_dim - self._dtype = input.dtype + self._workspaces[key] = workspace logger.info( f'FlashInfer all-reduce workspace initialized: rank={rank}, world_size={self._world_size}, ' f'max_token_num={max_token_num}, hidden_dim={hidden_dim}') - return self._workspace + return workspace def all_reduce_(self, input: torch.Tensor) -> bool: """All-reduce ``input`` in place, returning whether it was handled.""" @@ -171,7 +170,7 @@ def fused_all_reduce_residual_rms_norm(self, weight: torch.Tensor, eps: float): """Fuse all-reduce, residual addition and RMSNorm when supported.""" - if weight.dtype != input.dtype: + if input.dtype not in (torch.float16, torch.bfloat16) or weight.dtype != input.dtype: return None input_2d = input.flatten(0, -2) residual_2d = residual.flatten(0, -2) @@ -204,6 +203,6 @@ def fused_all_reduce_residual_rms_norm(self, def close(self): """Release the group-bound FlashInfer workspace.""" - if self._workspace is not None: - self._workspace.destroy() - self._workspace = None + for workspace in self._workspaces.values(): + workspace.destroy() + self._workspaces.clear() diff --git a/lmdeploy/pytorch/backends/cuda/hc_prepost.py b/lmdeploy/pytorch/backends/cuda/hc_prepost.py index d3b8956b5a..c9b81070c7 100644 --- a/lmdeploy/pytorch/backends/cuda/hc_prepost.py +++ b/lmdeploy/pytorch/backends/cuda/hc_prepost.py @@ -2,7 +2,7 @@ import torch from lmdeploy.pytorch.backends.hc_prepost import HCPrePostImpl -from lmdeploy.pytorch.kernels.cuda.dsv4.hc_prepost import hc_post_expand, hc_pre_reduce +from lmdeploy.pytorch.kernels.cuda.dsv4.hc_prepost import hc_post_expand, hc_pre_reduce, hc_pre_reduce_norm class TritonHCPrePostImpl(HCPrePostImpl): @@ -19,11 +19,16 @@ def pre( hc_scale: torch.Tensor, hc_base: torch.Tensor, out_dtype: torch.dtype, + norm_weight: torch.Tensor | None = None, + norm_eps: float = 1e-6, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: from lmdeploy.pytorch.kernels.cuda.dsv4.hc_split_sinkhorn import hc_split_sinkhorn pre, post, comb = hc_split_sinkhorn( mixes, hc_scale, hc_base, self.hc_mult, self.sinkhorn_iters, self.eps) - y = self.pre_reduce(x, pre, out_dtype) + if norm_weight is None: + y = self.pre_reduce(x, pre, out_dtype) + else: + y = hc_pre_reduce_norm(x, pre, self.hc_mult, norm_weight, norm_eps, out_dtype) return y, post, comb def pre_reduce(self, x: torch.Tensor, pre: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: diff --git a/lmdeploy/pytorch/backends/cuda/kda.py b/lmdeploy/pytorch/backends/cuda/kda.py index 7c94a4b07e..0ad3475a8b 100644 --- a/lmdeploy/pytorch/backends/cuda/kda.py +++ b/lmdeploy/pytorch/backends/cuda/kda.py @@ -135,7 +135,7 @@ def _forward_spec_decode(self, mixed_qkv, raw_gate, raw_beta, conv_state, raise ValueError('KDA verification exceeds the configured state ring.') history = metadata.cache_seqlens signed_ids = torch.where(metadata.valid_state, ids, -1) - values = mixed_qkv.reshape(batch, steps, -1).transpose(1, 2).contiguous() + values = mixed_qkv.reshape(batch, steps, -1).transpose(1, 2) weight = kwargs['conv_weight'] if weight.ndim == 3: if weight.size(1) != 1: diff --git a/lmdeploy/pytorch/backends/cuda/kpool.py b/lmdeploy/pytorch/backends/cuda/kpool.py index 6e48ec2c9f..87067dd78a 100644 --- a/lmdeploy/pytorch/backends/cuda/kpool.py +++ b/lmdeploy/pytorch/backends/cuda/kpool.py @@ -10,7 +10,13 @@ from lmdeploy.pytorch.kernels.cuda.fill_kv_cache import fill_indexed_key_cache from lmdeploy.pytorch.kernels.cuda.flatten_kv_cache import flatten_kv_cache -from lmdeploy.pytorch.kernels.cuda.kpool import compress_kpool, kpool_prefill_metadata, partition_kpool +from lmdeploy.pytorch.kernels.cuda.kpool import ( + compress_kpool, + kpool_prefill_metadata, + partition_kpool, + rotate_kpool_query, + update_kpool, +) from lmdeploy.pytorch.kernels.cuda.sparse_index_topk import ( is_sparse_index_topk_supported, sparse_index_topk, @@ -25,8 +31,31 @@ from .gated_delta_rule import _state_scatter # Fuse the existing integer index expansion instead of materializing its -# [tokens, topk] masks and int64 temporaries separately during batched prefill. -_expand_prefill_groups = torch.compile(kpool_expand_selected_groups, dynamic=True, fullgraph=True) +# [tokens, topk] masks and int64 temporaries separately during prefill or decode. +kpool_expand_groups_cuda = torch.compile(kpool_expand_selected_groups, dynamic=True, fullgraph=True) +kpool_rotate_query_cuda = rotate_kpool_query + + +def kpool_dense_indices_cuda(q_seqlens, kv_seqlens, rows, pool_size, topk): + """Reuse causal expansion when every prefill group fits in the budget.""" + _, _, seq, lengths, _, _ = kpool_prefill_metadata(q_seqlens, kv_seqlens, rows, pool_size) + groups = torch.arange(topk // pool_size, device=q_seqlens.device, dtype=torch.int32) + return kpool_expand_groups_cuda(groups.expand(rows, -1), lengths, pool_size, topk, seq_lens=seq) + + +def kpool_decode_update_cuda(keys, scores, tail_keys, tail_scores, state_ids, + history_lengths, packed_cache, block_offsets, + ape, pool_size, round_scale): + """Batch verification updates and reuse indexed cache writes.""" + closed_keys, closed_scores, groups, valid = update_kpool( + keys, scores, tail_keys, tail_scores, state_ids, history_lengths, pool_size) + compress = (compress_kpool if torch.cuda.get_device_capability(keys.device)[0] >= 9 + else kpool_compress_quantize_cuda) + values, scales = compress( + closed_keys, closed_scores, ape, mode='decode', round_scale=round_scale) + cache_keys, cache_scales = kpool_packed_cache_views(packed_cache, keys.size(-1)) + fill_indexed_key_cache(values, scales, groups, valid, block_offsets, + cache_keys, cache_scales, page_step=pool_size) def kpool_prefill_update_cuda(keys, scores, tail_keys, tail_scores, state_ids, @@ -118,7 +147,7 @@ def kpool_select_prefill_cuda(query_fp8, query_weight, packed_cache, # suffix using each query's causal length, including zero-length rows. selected = kpool_select_groups_cuda( logits, lengths, group_topk=topk // pool_size, max_group_length=max_groups) - return _expand_prefill_groups(selected, lengths, pool_size, topk, seq_lens=seq) + return kpool_expand_groups_cuda(selected, lengths, pool_size, topk, seq_lens=seq) def kpool_select_groups_cuda( diff --git a/lmdeploy/pytorch/backends/hc_prepost.py b/lmdeploy/pytorch/backends/hc_prepost.py index fe68a598fd..ea0cd1c6ae 100644 --- a/lmdeploy/pytorch/backends/hc_prepost.py +++ b/lmdeploy/pytorch/backends/hc_prepost.py @@ -18,9 +18,11 @@ def pre( hc_scale: torch.Tensor, hc_base: torch.Tensor, out_dtype: torch.dtype, + norm_weight: torch.Tensor | None = None, + norm_eps: float = 1e-6, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: """Run sinkhorn and reduce HC states from ``[..., hc, dim]`` to ``[..., - dim]``.""" + dim]``, optionally followed by RMSNorm after the output-dtype cast.""" raise NotImplementedError @abstractmethod diff --git a/lmdeploy/pytorch/kernels/cuda/causal_conv1d.py b/lmdeploy/pytorch/kernels/cuda/causal_conv1d.py index 0dbb7e8e46..5b88bf0106 100644 --- a/lmdeploy/pytorch/kernels/cuda/causal_conv1d.py +++ b/lmdeploy/pytorch/kernels/cuda/causal_conv1d.py @@ -217,7 +217,8 @@ def causal_conv1d_fn( }, ) def causal_conv1d_update_fwd(hidden_size: int, seqlen: int, state_len: int, width: int, has_bias: bool, activation: str | None, dtype, conv_stride: tuple[int, int, int], is_circular_buffer: bool, - has_state_indices: bool, num_warps: int, weight_dtype=None, bias_dtype=None): + has_state_indices: bool, num_warps: int, weight_dtype=None, bias_dtype=None, + x_stride=None): """TileLang kernel for causal convolution forward pass. Each thread processes one output position for all channels sequentially. @@ -231,11 +232,13 @@ def causal_conv1d_update_fwd(hidden_size: int, seqlen: int, state_len: int, widt batch = T.dynamic('batch') conv_batch = T.dynamic('conv_batch') conv_batch_stride = T.dynamic('conv_batch_stride') + if x_stride is None: + x_stride = (hidden_size * seqlen, seqlen, 1) update_idx_base = -(width - 1) @T.prim_func def causal_conv1d_update_main( - X: T.Tensor((batch, hidden_size, seqlen), dtype=dtype), + X: T.StridedTensor((batch, hidden_size, seqlen), dtype=dtype, strides=x_stride), Conv_State: T.StridedTensor((conv_batch, hidden_size, state_len), dtype=dtype, strides=(conv_batch_stride, conv_stride[1], conv_stride[2])), @@ -372,7 +375,8 @@ def causal_conv1d_update(x, has_state_indices=conv_state_indices is not None, num_warps=num_warps, weight_dtype=weight.dtype, - bias_dtype=bias.dtype if bias is not None else x.dtype) + bias_dtype=bias.dtype if bias is not None else x.dtype, + x_stride=x.stride()) kernel(x, conv_state, weight, bias, out, cache_seqlens, conv_state_indices) diff --git a/lmdeploy/pytorch/kernels/cuda/dsv4/hc_prepost.py b/lmdeploy/pytorch/kernels/cuda/dsv4/hc_prepost.py index 626669165e..9f88dba564 100644 --- a/lmdeploy/pytorch/kernels/cuda/dsv4/hc_prepost.py +++ b/lmdeploy/pytorch/kernels/cuda/dsv4/hc_prepost.py @@ -3,6 +3,8 @@ import triton import triton.language as tl +from ..rms_norm import _compute_rms_norm + def _get_block_d(dim: int) -> int: if dim <= 64: @@ -46,6 +48,52 @@ def _hc_pre_reduce_kernel( tl.store(out_ptr + row_id * out_stride_n + offs_d * out_stride_d, acc, mask=mask) +@triton.jit +def _hc_pre_reduce_norm_kernel( + x_ptr, pre_ptr, weight_ptr, out_ptr, + x_stride_n, x_stride_h, x_stride_d, + pre_stride_n, pre_stride_h, + dim: tl.constexpr, hc_mult: tl.constexpr, eps: tl.constexpr, + BLOCK_D: tl.constexpr, +): + row_id = tl.program_id(0) + offs_d = tl.arange(0, BLOCK_D) + mask = offs_d < dim + acc = tl.zeros((BLOCK_D,), dtype=tl.float32) + for hc_id in range(hc_mult): + pre = tl.load(pre_ptr + row_id * pre_stride_n + hc_id * pre_stride_h).to(tl.float32) + x = tl.load(x_ptr + row_id * x_stride_n + hc_id * x_stride_h + offs_d * x_stride_d, + mask=mask, other=0.0).to(tl.float32) + acc += pre * x + + # Preserve the separate pre-reduction's store/cast before RMSNorm. + reduced = acc.to(out_ptr.dtype.element_ty) + weight = tl.load(weight_ptr + offs_d, mask=mask, other=0.0) + out = _compute_rms_norm(reduced, weight, eps, dim) + tl.store(out_ptr + row_id * dim + offs_d, out, mask=mask) + + +def hc_pre_reduce_norm(x: torch.Tensor, pre: torch.Tensor, hc_mult: int, + weight: torch.Tensor, eps: float, + out_dtype: torch.dtype) -> torch.Tensor: + """Fuse HC reduction and RMSNorm while retaining the intermediate cast.""" + dim = x.size(-1) + assert weight.shape == (dim,) and weight.is_contiguous() + out_shape = (*x.shape[:-2], dim) + out = torch.empty(out_shape, device=x.device, dtype=out_dtype) + if x.numel() == 0: + return out + x = x.reshape(-1, hc_mult, dim) + pre = pre.reshape(-1, hc_mult) + block_d = triton.next_power_of_2(dim) + _hc_pre_reduce_norm_kernel[(x.size(0),)]( + x, pre, weight, out, *x.stride(), *pre.stride(), + dim, hc_mult, eps, block_d, + num_warps=min(triton.cdiv(block_d, 2048), 4), + ) + return out + + @triton.jit def _hc_post_expand_kernel( x_ptr, @@ -154,6 +202,10 @@ def hc_post_expand( out = out.reshape(-1, hc_mult, dim) n_rows = x.size(0) block_d = _get_block_d(dim) + # Large prefills have enough independent rows; use wider memory tiles to + # reduce CTA count while retaining the same per-element accumulation. + if hc_mult == 4 and n_rows >= 128: + block_d = min(triton.next_power_of_2(dim), 1024) grid = (n_rows * hc_mult, triton.cdiv(dim, block_d)) _hc_post_expand_kernel[grid]( x, diff --git a/lmdeploy/pytorch/kernels/cuda/gated_delta_rule.py b/lmdeploy/pytorch/kernels/cuda/gated_delta_rule.py index c34596f806..f83dd31c24 100644 --- a/lmdeploy/pytorch/kernels/cuda/gated_delta_rule.py +++ b/lmdeploy/pytorch/kernels/cuda/gated_delta_rule.py @@ -308,7 +308,9 @@ def fused_recurrent_gated_delta_rule_fwd(SEQLEN, desired = T.ceildiv(V, num_warps) v_per_warp = T.ceildiv(min(desired, max_v_per_warp), min_v_per_warp) * min_v_per_warp v_per_warp = max(v_per_warp, min_v_per_warp) - target_v_per_cta = V + # Channelwise KDA has independent V tiles. Distribute them across + # CTAs instead of serial state waves, especially for small batches. + target_v_per_cta = v_per_warp * num_warps if channelwise_g and K == V == 128 else V else: target_v_per_cta = max(V, v_per_warp * num_warps * 2) diff --git a/lmdeploy/pytorch/kernels/cuda/kpool.py b/lmdeploy/pytorch/kernels/cuda/kpool.py index 1238ac41f3..cabea0abeb 100644 --- a/lmdeploy/pytorch/kernels/cuda/kpool.py +++ b/lmdeploy/pytorch/kernels/cuda/kpool.py @@ -5,6 +5,88 @@ import triton.language.extra.cuda.libdevice as libdevice +@triton.jit +def _update_kpool_kernel( + Keys, Scores, TailKeys, TailScores, StateIds, History, + ClosedKeys, ClosedScores, GroupIds, Valid, + BATCH: tl.constexpr, STEPS: tl.constexpr, STATES: tl.constexpr, RING: tl.constexpr, + POOL: tl.constexpr, WIDTH: tl.constexpr, BLOCK_D: tl.constexpr, + stride_kb: tl.constexpr, stride_kt: tl.constexpr, stride_kd: tl.constexpr, + stride_sb: tl.constexpr, stride_st: tl.constexpr, stride_sd: tl.constexpr, + stride_tkb: tl.constexpr, stride_tkr: tl.constexpr, + stride_tsb: tl.constexpr, stride_tsr: tl.constexpr, +): + request = tl.program_id(0) + state_id = tl.load(StateIds + request).to(tl.int64) + valid = (state_id >= 0) & (state_id < STATES) + history = tl.maximum(tl.load(History + request).to(tl.int64), 0) + slot = tl.arange(0, POOL)[:, None] + d = tl.arange(0, BLOCK_D)[None, :] + mask = valid & (d < WIDTH) + tail_k = tl.load(TailKeys + state_id * stride_tkb + history % RING * stride_tkr + slot * WIDTH + d, + mask & (slot < history % POOL), other=0) + tail_s = tl.load(TailScores + state_id * stride_tsb + history % RING * stride_tsr + slot * WIDTH + d, + mask & (slot < history % POOL), other=0) + # A request owns its state row. Keep time sequential in registers, and + # persist every checkpoint so rejection can resume at any accepted prefix. + for step in range(STEPS): + position = (history + step) % POOL + key = tl.load(Keys + request * stride_kb + step * stride_kt + d * stride_kd, mask, other=0) + score = tl.load(Scores + request * stride_sb + step * stride_st + d * stride_sd, mask, other=0) + tail_k = tl.where(slot == position, key, tail_k) + tail_s = tl.where(slot == position, score, tail_s) + row = step * BATCH + request + offset = (row * POOL + slot) * WIDTH + d + tl.store(ClosedKeys + offset, tail_k, d < WIDTH) + tl.store(ClosedScores + offset, tail_s, d < WIDTH) + close = position == POOL - 1 + tl.store(GroupIds + row, (history + step) // POOL) + tl.store(Valid + row, valid & close) + tail_k = tl.where(close, 0, tail_k) + tail_s = tl.where(close, 0, tail_s) + checkpoint = (history + step + 1) % RING + tl.store(TailKeys + state_id * stride_tkb + checkpoint * stride_tkr + slot * WIDTH + d, tail_k, mask) + tl.store(TailScores + state_id * stride_tsb + checkpoint * stride_tsr + slot * WIDTH + d, tail_s, mask) + + +def update_kpool(keys, scores, tail_keys, tail_scores, state_ids, history_lengths, pool_size): + """Update decode/verify tails in place, emitting step-major closed pools. + + Inputs use [batch, steps, width]. State uses [states, pool, width] for AR or [states, ring, pool, width] for + verification. Live state ids must be unique; invalid ids never write state. All output shapes depend only on input + shapes. + """ + if keys.ndim != 3 or scores.shape != keys.shape: + raise ValueError('Expected matching [batch, steps, width] keys and scores.') + batch, steps, width = keys.shape + if not batch or not steps or pool_size <= 1 or pool_size & (pool_size - 1): + raise ValueError('Expected a nonempty batch/sequence and a power-of-two pool size.') + if state_ids.shape != (batch,) or history_lengths.shape != (batch,): + raise ValueError('Expected one state id and history length per request.') + if tail_keys.shape != tail_scores.shape or tail_keys.ndim not in (3, 4): + raise ValueError('Expected matching AR or verification tail caches.') + if tail_keys.ndim == 3: + tail_keys, tail_scores = tail_keys.unsqueeze(1), tail_scores.unsqueeze(1) + elif steps > tail_keys.size(1): + raise ValueError('The checkpoint ring must hold every verification step.') + if tail_keys.shape[2:] != (pool_size, width): + raise ValueError('Tail cache geometry does not match decode inputs.') + if tail_keys.stride()[2:] != (width, 1) or tail_scores.stride()[2:] != (width, 1): + raise ValueError('Tail cache rows must be contiguous.') + if keys.dtype != tail_keys.dtype or scores.dtype != tail_scores.dtype: + raise ValueError('Keys and scores must match their respective tail cache dtypes.') + closed_keys = keys.new_empty((steps * batch, pool_size, width)) + closed_scores = scores.new_empty((steps * batch, pool_size, width)) + groups = history_lengths.new_empty(steps * batch) + valid = torch.empty(steps * batch, device=keys.device, dtype=torch.bool) + _update_kpool_kernel[(batch,)]( + keys, scores, tail_keys, tail_scores, state_ids.contiguous(), history_lengths.contiguous(), + closed_keys, closed_scores, groups, valid, + batch, steps, tail_keys.size(0), tail_keys.size(1), pool_size, width, triton.next_power_of_2(width), + *keys.stride(), *scores.stride(), *tail_keys.stride()[:2], *tail_scores.stride()[:2], num_warps=4) + return closed_keys, closed_scores, groups, valid + + @triton.jit def _prefill_offsets_kernel(Q, KV, Counts, Starts, QEnds, BATCH: tl.constexpr, POOL: tl.constexpr, BLOCK: tl.constexpr): @@ -160,6 +242,43 @@ def partition_kpool(keys, scores, tail_keys, tail_scores, state_ids, q_seqlens, return closed_keys, closed_scores, groups, requests, valid, next_keys, next_scores +@triton.jit +def _normalized_hadamard(value, WIDTH: tl.constexpr, LEVELS: tl.constexpr): + d = tl.arange(0, WIDTH) + if len(value.shape) == 2: + d = d[None, :] + for level in tl.static_range(LEVELS): + stride = 1 << level + other = tl.gather(value, tl.broadcast_to(d ^ stride, value.shape), axis=len(value.shape) - 1) + value = tl.where((d & stride) == 0, value + other, other - value) + return value * (WIDTH**-0.5) + + +@triton.jit +def _rotate_kpool_query_kernel(Query, Out, ROWS: tl.constexpr, WIDTH: tl.constexpr, + LEVELS: tl.constexpr, BLOCK_M: tl.constexpr): + rows = tl.program_id(0) * BLOCK_M + tl.arange(0, BLOCK_M) + offsets = rows[:, None] * WIDTH + tl.arange(0, WIDTH)[None, :] + value = tl.load(Query + offsets, rows[:, None] < ROWS, other=0).to(tl.float32) + value = _normalized_hadamard(value, WIDTH, LEVELS) + tl.store(Out + offsets, value, rows[:, None] < ROWS) + + +def rotate_kpool_query(query: torch.Tensor) -> torch.Tensor: + """Reuse the compression butterfly, retaining FP32 arithmetic and the + output cast.""" + width = query.size(-1) + if width <= 0 or width & (width - 1): + raise ValueError('Query width must be a positive power of two.') + out = torch.empty_like(query, memory_format=torch.contiguous_format) + rows = query.numel() // width + if rows: + _rotate_kpool_query_kernel[(triton.cdiv(rows, 4),)]( + query.contiguous(), out, rows, width, width.bit_length() - 1, 4, + num_warps=4, enable_fp_fusion=False) + return out + + @triton.jit def _compress_kpool_kernel(K, S, A, O, Scale, WIDTH: tl.constexpr, POOL: tl.constexpr, ONLINE: tl.constexpr, ROUND: tl.constexpr, LEVELS: tl.constexpr): @@ -187,11 +306,7 @@ def _compress_kpool_kernel(K, S, A, O, Scale, WIDTH: tl.constexpr, POOL: tl.cons denominator = denominator + probability accumulator = accumulator + key * probability value = tl.div_rn(accumulator, denominator).to(tl.bfloat16).to(tl.float32) - for level in tl.static_range(LEVELS): - stride = 1 << level - other = tl.gather(value, d ^ stride, axis=0) - value = tl.where((d & stride) == 0, value + other, other - value) - value = (value * (WIDTH**-0.5)).to(tl.bfloat16).to(tl.float32) + value = _normalized_hadamard(value, WIDTH, LEVELS).to(tl.bfloat16).to(tl.float32) scale = tl.maximum(tl.max(tl.abs(value), 0), 1e-4) * (1.0 / 448.0) if ROUND: scale = libdevice.exp2(libdevice.ceil(libdevice.log2(scale))) diff --git a/lmdeploy/pytorch/models/glm5_next.py b/lmdeploy/pytorch/models/glm5_next.py index ae9d2b6580..a9f6b76ebb 100644 --- a/lmdeploy/pytorch/models/glm5_next.py +++ b/lmdeploy/pytorch/models/glm5_next.py @@ -16,7 +16,11 @@ from lmdeploy.pytorch.backends.cuda.attention.sparse_mla import FlashMLASparseImpl from lmdeploy.pytorch.backends.cuda.kpool import ( kpool_compress_quantize_cuda, + kpool_decode_update_cuda, + kpool_dense_indices_cuda, + kpool_expand_groups_cuda, kpool_prefill_update_cuda, + kpool_rotate_query_cuda, kpool_score_contiguous_cuda, kpool_score_paged_cuda, kpool_select_groups_cuda, @@ -29,7 +33,7 @@ GLM5_KPOOL_TAIL_K_STATE, GLM5_KPOOL_TAIL_SCORE_STATE, ) -from lmdeploy.pytorch.distributed import get_dist_manager, get_tp_world_rank +from lmdeploy.pytorch.distributed import get_dist_group, get_dist_manager, get_tp_world_rank from lmdeploy.pytorch.engine.cache_engine.schema import BlockCacheRequest from lmdeploy.pytorch.model_inputs import StepContext, StepContextManager, get_step_ctx_manager from lmdeploy.pytorch.nn import ( @@ -544,7 +548,7 @@ def forward( if self._fp32_tp_reduce: output_dtype = out.dtype out = out.float() - dist.all_reduce(out, group=self.experts.tp_group) + get_dist_group('moe').all_reduce_(out) out = out.to(output_dtype) return out @@ -879,6 +883,22 @@ def _update_kpool_cache( tail_k_state, tail_score_state = tail_state history_lengths = attn_metadata.kv_seqlens - attn_metadata.q_seqlens + indexer_k_cache = self.indexer.get_block_cache() + key = self.indexer.project_key(hidden_states)[0] + score = self.indexer.project_compress_score(hidden_states)[0] + if attn_metadata.is_decoding and key.is_cuda: + batch_size = state_ids.numel() + if not batch_size or key.size(0) % batch_size: + raise RuntimeError('KPool decode rows must be divisible by request count.') + steps = key.size(0) // batch_size + kpool_decode_update_cuda( + key.unflatten(0, (batch_size, steps)), score.unflatten(0, (batch_size, steps)), + tail_k_state, tail_score_state, state_ids, history_lengths, + indexer_k_cache, attn_metadata.block_offsets, + self.indexer.index_kpool_compress_ape, self.index_kpool, + self.indexer.scale_fmt is not None) + return indexer_k_cache + ring_states = None if tail_k_state.ndim == 4: ring_states = tail_state @@ -901,9 +921,6 @@ def save_ring(lengths): state[request_ids, slots] = torch.where( valid_requests[:, None, None], value, previous) - indexer_k_cache = self.indexer.get_block_cache() - key = self.indexer.project_key(hidden_states)[0] - score = self.indexer.project_compress_score(hidden_states)[0] if attn_metadata.is_decoding: batch_size = state_ids.numel() if key.size(0) % batch_size: @@ -1025,31 +1042,36 @@ def _select_kpool_indices( indexer_k_cache: torch.Tensor, attn_metadata: Any, ) -> torch.Tensor: - """Score/select on attention-TP rank 0, then broadcast logical ids.""" + """Score/select on rank 0; expand decode group ids on each TP rank.""" + if (hidden_states.is_cuda and not attn_metadata.is_decoding + and attn_metadata.max_kv_seqlen <= self.index_topk): + # The shared Top-K returns ascending ids when every group fits. + # MTP still needs these seed rows, but no scores or TP broadcast. + return kpool_dense_indices_cuda( + attn_metadata.q_seqlens, attn_metadata.kv_seqlens, + hidden_states.size(1), self.index_kpool, self.index_topk) dist_ctx = get_dist_manager().current_context() tp_group = dist_ctx.attn_tp_group is_owner = tp_group.rank == 0 total_rows = hidden_states.size(1) output_width = self.index_topk + self.index_kpool - 1 + if attn_metadata.is_decoding: + batch_size = attn_metadata.kv_seqlens.numel() + steps = total_rows // batch_size + history = attn_metadata.kv_seqlens - attn_metadata.q_seqlens + step_ids = torch.arange(1, steps + 1, device=hidden_states.device) + seq_lens = (history[:, None] + step_ids).flatten().to(torch.int64) + group_lengths = torch.div(seq_lens, self.index_kpool, rounding_mode='floor') if is_owner: query = self.indexer.project_query(q_lora)[0] - query = kpool_rotate_query(query) + rotate = kpool_rotate_query_cuda if query.is_cuda else kpool_rotate_query + query = rotate(query) query_fp8, query_scale = self.indexer.quantize_fp8(query) head_gate = self.indexer.project_head_gate(hidden_states)[0] query_weight = (head_gate * query_scale.squeeze(-1) * self.indexer.softmax_scale) if attn_metadata.is_decoding: - batch_size = attn_metadata.kv_seqlens.numel() - steps = total_rows // batch_size - history = attn_metadata.kv_seqlens - attn_metadata.q_seqlens - step_ids = torch.arange(1, steps + 1, device=query_fp8.device) - seq_lens = (history[:, None] + step_ids).flatten().to(torch.int64) - group_lengths = torch.div( - seq_lens, - self.index_kpool, - rounding_mode='floor', - ) pooled_block_offsets = kpool_pooled_block_offsets( attn_metadata.block_offsets.repeat_interleave(steps, dim=0), self.index_kpool, @@ -1066,13 +1088,7 @@ def _select_kpool_indices( group_lengths, group_topk=self.index_topk // self.index_kpool, ) - logical_indices = kpool_expand_selected_groups( - selected_groups, - group_lengths, - self.index_kpool, - self.index_topk, - seq_lens=seq_lens, - ) + logical_indices = selected_groups else: logical_indices = self._select_kpool_indices_prefill( query_fp8, @@ -1083,7 +1099,7 @@ def _select_kpool_indices( else: logical_indices = torch.empty( total_rows, - output_width, + self.index_topk // self.index_kpool if attn_metadata.is_decoding else output_width, dtype=torch.int32, device=hidden_states.device, ) @@ -1092,6 +1108,10 @@ def _select_kpool_indices( group = tp_group.gpu_group source_rank = dist_ctx.rank - tp_group.rank dist.broadcast(logical_indices, src=source_rank, group=group) + if attn_metadata.is_decoding: + expand = kpool_expand_groups_cuda if logical_indices.is_cuda else kpool_expand_selected_groups + logical_indices = expand(logical_indices, group_lengths, + self.index_kpool, self.index_topk, seq_lens=seq_lens) return logical_indices def _select_kpool_indices_prefill( @@ -1409,7 +1429,7 @@ def __init__(self, def _hc_pre(self, hidden_states: torch.Tensor, fn: torch.Tensor, scale: torch.Tensor, base: torch.Tensor, norm: RMSNorm): return self.hc_prepost.pre( - hidden_states, fn, scale, base, norm.eps) + hidden_states, fn, scale, base, norm.eps, norm_weight=norm.weight) def forward(self, hidden_states: torch.Tensor, past_key_value: Sequence[torch.Tensor], attn_metadata: Any, @@ -1424,7 +1444,6 @@ def forward(self, hidden_states: torch.Tensor, self.hc_attn_base, self.input_layernorm, ) - hidden_states = self.input_layernorm(hidden_states) if self.is_linear_attention: hidden_states = self.self_attn(hidden_states, past_key_value=past_key_value, @@ -1446,7 +1465,6 @@ def forward(self, hidden_states: torch.Tensor, self.hc_ffn_base, self.post_attention_layernorm, ) - hidden_states = self.post_attention_layernorm(hidden_states) hidden_states = self.mlp(hidden_states) return self.hc_prepost.post_expand(hidden_states, residual, post, comb) diff --git a/lmdeploy/pytorch/nn/hc_prepost.py b/lmdeploy/pytorch/nn/hc_prepost.py index 3f7d1d82c8..5b31db8df7 100644 --- a/lmdeploy/pytorch/nn/hc_prepost.py +++ b/lmdeploy/pytorch/nn/hc_prepost.py @@ -26,6 +26,7 @@ def pre( hc_scale: torch.Tensor, hc_base: torch.Tensor, norm_eps: float, + norm_weight: torch.Tensor | None = None, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: from lmdeploy.pytorch.nn.norm import rms_scale hidden_states, dtype = x, x.dtype @@ -37,7 +38,8 @@ def pre( else: mixes = F.linear(x, hc_fn) mixes = rms_scale(mixes, x, eps=norm_eps) - return self.impl.pre(hidden_states, mixes, hc_scale, hc_base, dtype) + return self.impl.pre(hidden_states, mixes, hc_scale, hc_base, dtype, + norm_weight=norm_weight, norm_eps=norm_eps) def pre_reduce(self, x: torch.Tensor, pre: torch.Tensor, out_dtype: torch.dtype) -> torch.Tensor: return self.impl.pre_reduce(x, pre, out_dtype) diff --git a/lmdeploy/pytorch/nn/linear/base.py b/lmdeploy/pytorch/nn/linear/base.py index 77c166ef3e..6dbd7b3ed9 100644 --- a/lmdeploy/pytorch/nn/linear/base.py +++ b/lmdeploy/pytorch/nn/linear/base.py @@ -214,7 +214,9 @@ def _forward_lora(self, x, tp_sizes: list[int] = None): output_dtype = out.dtype if self.tp_reduce_dtype is not None: out = out.to(self.tp_reduce_dtype) - dist.all_reduce(out, group=self.tp_group) + get_dist_group(self.layer_type).all_reduce_(out) + else: + dist.all_reduce(out, group=self.tp_group) out = out.to(output_dtype) return out From e2001908ca2dad8bba70cbb4a8eae525be1cba46 Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Mon, 28 Sep 2026 07:49:17 +0000 Subject: [PATCH 28/39] fix(pytorch): accumulate DeepEP local expert reduction in FP32 Use an FP32 accumulator in ep_gather independently of the output dtype, matching vLLM's DeepGEMM unpermute-and-reduce. Store the result directly as BF16 for DeepEP combine, removing the FP32 output buffer and separate cast from GLM's normal path. Remove the private fp32_acc plumbing from the blocked-FP8 DeepEP builder and normal execution paths. Preserve activation callbacks, FP32 scaled routing weights, and the existing low-latency combine implementation. The shared gather now uses FP32 for other callers, including BF16 experts. Validation: 16 H200 kernel cases matched vLLM 606d124b and the previous FP32-output-then-cast path exactly; covered BF16/FP16, top-k 1/8, scales 1/2.5, missing experts, row strides, and more than 1024 tokens. CUDA graph replay and cancellation regression passed. Ruff 0.15.4, Python compile, and git diff --check passed. Full-model/multi-rank inference not rerun. --- lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py | 11 +++-------- lmdeploy/pytorch/kernels/cuda/moe/ep.py | 8 ++------ lmdeploy/pytorch/kernels/cuda/moe/ep_fp8.py | 11 ++++------- 3 files changed, 9 insertions(+), 21 deletions(-) diff --git a/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py b/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py index bba12e1821..bc22404c41 100644 --- a/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py +++ b/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py @@ -46,7 +46,6 @@ def __init__( num_max_dispatch_tokens_per_rank: int = 128, chunk_size: int | None = 32 * 1024, expert_alignment: int = 128, - fp32_acc: bool = False, ): self.layer_index = layer_index self.top_k = top_k @@ -56,7 +55,6 @@ def __init__( self.out_dtype = out_dtype self.fp8_dtype = fp8_dtype self.scale_fmt = scale_fmt - self.fp32_acc = fp32_acc self.token_dispatcher = DeepEPTokenDispatcherNormal( group=ep_group, num_experts=num_experts, @@ -91,7 +89,7 @@ def forward( ) out_states = fused_moe_v3_fp8(x, recv_topk_ids, recv_topk_weights, (up_weights, up_scale), (down_weights, down_scale), recv_tokens_per_expert, - act_func=act_func, scale_fmt=self.scale_fmt, fp32_acc=self.fp32_acc) + act_func=act_func, scale_fmt=self.scale_fmt) return self.token_dispatcher.combine(out_states) def capture(self): @@ -122,7 +120,7 @@ def release(self): def fusedmoe_forward(self, state, up_weight, up_scale, down_weight, down_scale, act_func=None): return fused_moe_v3_fp8(state['recv_hidden_states'], state['recv_topk_idx'], state['recv_topk_weights'], (up_weight, up_scale), (down_weight, down_scale), state['recv_tokens_per_expert'], - act_func=act_func, scale_fmt=self.scale_fmt, fp32_acc=self.fp32_acc) + act_func=act_func, scale_fmt=self.scale_fmt) def per_token_group_quant_fp8(self, x: torch.Tensor, @@ -274,7 +272,6 @@ def _build_deepep_moe( num_max_dispatch_tokens_per_rank: int = 128, chunk_size: int | None = 32 * 1024, expert_alignment: int = 128, - fp32_acc: bool = False, ): if low_latency_mode: return FusedMoELowLatency(ep_size=ep_size, @@ -298,8 +295,7 @@ def _build_deepep_moe( scale_fmt=scale_fmt, num_max_dispatch_tokens_per_rank=num_max_dispatch_tokens_per_rank, chunk_size=chunk_size, - expert_alignment=expert_alignment, - fp32_acc=fp32_acc) + expert_alignment=expert_alignment) class TritonFusedMoEBlockedF8Impl(FusedMoEBlockedF8Impl): @@ -563,7 +559,6 @@ def fusedmoe_build(self, low_latency_mode: bool = False): scale_fmt=self.scale_fmt, layer_idx=self.layer_idx, num_max_dispatch_tokens_per_rank=self.num_max_dispatch_tokens_per_rank, - fp32_acc=self.output_scale != 1.0, chunk_size=16 * 1024) return deepep_moe diff --git a/lmdeploy/pytorch/kernels/cuda/moe/ep.py b/lmdeploy/pytorch/kernels/cuda/moe/ep.py index ed51f840a8..53ee8f5a19 100644 --- a/lmdeploy/pytorch/kernels/cuda/moe/ep.py +++ b/lmdeploy/pytorch/kernels/cuda/moe/ep.py @@ -146,20 +146,16 @@ def _fwd_kernel_ep_gather( cur_block = tl.program_id(0) start_cur_token = tl.program_id(1) grid_num = tl.num_programs(1) - # align with xtuner rl - compute_dtype = output_tensor.dtype.element_ty - # compute_dtype = tl.float32 - for cur_token in range(start_cur_token, total_token_num, grid_num): off_d = tl.arange(0, BLOCK_D) - accumulator = tl.zeros([BLOCK_D], dtype=compute_dtype) + accumulator = tl.zeros([BLOCK_D], dtype=tl.float32) for topk_index in range(0, topk_num): expert_id = tl.load(recv_topk_ids + cur_token * recv_topk_ids_stride0 + topk_index) if expert_id >= 0: source_token_index = tl.load(input_index + cur_token * input_index_stride0 + topk_index) acc_weight = tl.load(recv_topk_weight + cur_token * recv_topk_weight_stride0 + topk_index) tmp = tl.load(input_tensor + source_token_index * input_tensor_stride0 + cur_block * BLOCK_D + off_d) - accumulator += tmp.to(compute_dtype) * acc_weight.to(compute_dtype) + accumulator += tmp.to(tl.float32) * acc_weight.to(tl.float32) tl.store( output_tensor + cur_token * output_tensor_stride0 + cur_block * BLOCK_D + off_d, accumulator.to(output_tensor.dtype.element_ty), diff --git a/lmdeploy/pytorch/kernels/cuda/moe/ep_fp8.py b/lmdeploy/pytorch/kernels/cuda/moe/ep_fp8.py index ab2f39c0b5..a092381d29 100644 --- a/lmdeploy/pytorch/kernels/cuda/moe/ep_fp8.py +++ b/lmdeploy/pytorch/kernels/cuda/moe/ep_fp8.py @@ -156,7 +156,6 @@ def fused_moe_v3_fp8( num_recv_tokens_per_expert: list[int] | None, act_func: Callable | None = None, scale_fmt: str | None = None, - fp32_acc: bool = False, ): hidden_states_fp8, hidden_states_scale = hidden_states_fp8 if num_recv_tokens_per_expert is None: @@ -168,9 +167,8 @@ def fused_moe_v3_fp8( m, k = hidden_states_fp8.size() n = w13_weight_fp8[0].size(1) block_size = k // hidden_states_scale.size(1) - # ep_gather already accumulates in its output dtype. Reuse that contract - # instead of introducing a second reduction kernel for FP32 accumulation. - gather_out = torch.empty_like(hidden_states_fp8, dtype=torch.float32 if fp32_acc else torch.bfloat16) + # ep_gather accumulates in FP32 and casts once when storing BF16 for DeepEP. + gather_out = torch.empty_like(hidden_states_fp8, dtype=torch.bfloat16) input_tensor = torch.empty((all_tokens, k), device=hidden_states_fp8.device, dtype=hidden_states_fp8.dtype) input_tensor_scale = torch.empty((all_tokens, k // block_size), device=hidden_states_fp8.device, @@ -200,6 +198,5 @@ def fused_moe_v3_fp8( down_output = torch.empty((all_tokens, k), device=gather_out.device, dtype=torch.bfloat16) _deepgemm_grouped_fp8_nt_contiguous((down_input_fp8, down_input_scale), w2_weight_fp8, down_output, m_indices) ep_gather(down_output, topk_idx, topk_weights, output_index, gather_out) - # DeepEP transports BF16 partial sums. Local expert reduction can be FP32, - # but this is not an all-FP32 cross-rank reduction contract. - return gather_out.to(torch.bfloat16) + # DeepEP transports BF16 partial sums; only the local reduction is FP32. + return gather_out From 8c167ebafe93602af085c8bc2ad3091c50454059 Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Mon, 28 Sep 2026 08:07:05 +0000 Subject: [PATCH 29/39] fix(pytorch): make DeepEP FP32 local accumulation opt-in Restore fp32_acc=False plumbing through the DeepEP Normal builder and both execution entry points. Pass the flag into ep_gather as a Triton constexpr so default callers retain output-dtype accumulation while GLM's scaled routed-expert path explicitly uses FP32. Keep the BF16 gather output and cast inside the kernel, avoiding the previous FP32 temporary and separate output cast. Preserve activation callbacks, routing-weight scaling and the low-latency combine path. Validation: 16 H200 cases matched legacy accumulation with the default and explicit False, and vLLM/previous GLM FP32 results with True. CUDA graph replay, cancellation, builder/sync/async parameter propagation, Ruff 0.15.4, Python compile and git diff --check passed. Full-model and multi-rank inference were not rerun. --- lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py | 11 ++++++++--- lmdeploy/pytorch/kernels/cuda/moe/ep.py | 8 ++++++-- lmdeploy/pytorch/kernels/cuda/moe/ep_fp8.py | 6 +++--- 3 files changed, 17 insertions(+), 8 deletions(-) diff --git a/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py b/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py index bc22404c41..bba12e1821 100644 --- a/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py +++ b/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py @@ -46,6 +46,7 @@ def __init__( num_max_dispatch_tokens_per_rank: int = 128, chunk_size: int | None = 32 * 1024, expert_alignment: int = 128, + fp32_acc: bool = False, ): self.layer_index = layer_index self.top_k = top_k @@ -55,6 +56,7 @@ def __init__( self.out_dtype = out_dtype self.fp8_dtype = fp8_dtype self.scale_fmt = scale_fmt + self.fp32_acc = fp32_acc self.token_dispatcher = DeepEPTokenDispatcherNormal( group=ep_group, num_experts=num_experts, @@ -89,7 +91,7 @@ def forward( ) out_states = fused_moe_v3_fp8(x, recv_topk_ids, recv_topk_weights, (up_weights, up_scale), (down_weights, down_scale), recv_tokens_per_expert, - act_func=act_func, scale_fmt=self.scale_fmt) + act_func=act_func, scale_fmt=self.scale_fmt, fp32_acc=self.fp32_acc) return self.token_dispatcher.combine(out_states) def capture(self): @@ -120,7 +122,7 @@ def release(self): def fusedmoe_forward(self, state, up_weight, up_scale, down_weight, down_scale, act_func=None): return fused_moe_v3_fp8(state['recv_hidden_states'], state['recv_topk_idx'], state['recv_topk_weights'], (up_weight, up_scale), (down_weight, down_scale), state['recv_tokens_per_expert'], - act_func=act_func, scale_fmt=self.scale_fmt) + act_func=act_func, scale_fmt=self.scale_fmt, fp32_acc=self.fp32_acc) def per_token_group_quant_fp8(self, x: torch.Tensor, @@ -272,6 +274,7 @@ def _build_deepep_moe( num_max_dispatch_tokens_per_rank: int = 128, chunk_size: int | None = 32 * 1024, expert_alignment: int = 128, + fp32_acc: bool = False, ): if low_latency_mode: return FusedMoELowLatency(ep_size=ep_size, @@ -295,7 +298,8 @@ def _build_deepep_moe( scale_fmt=scale_fmt, num_max_dispatch_tokens_per_rank=num_max_dispatch_tokens_per_rank, chunk_size=chunk_size, - expert_alignment=expert_alignment) + expert_alignment=expert_alignment, + fp32_acc=fp32_acc) class TritonFusedMoEBlockedF8Impl(FusedMoEBlockedF8Impl): @@ -559,6 +563,7 @@ def fusedmoe_build(self, low_latency_mode: bool = False): scale_fmt=self.scale_fmt, layer_idx=self.layer_idx, num_max_dispatch_tokens_per_rank=self.num_max_dispatch_tokens_per_rank, + fp32_acc=self.output_scale != 1.0, chunk_size=16 * 1024) return deepep_moe diff --git a/lmdeploy/pytorch/kernels/cuda/moe/ep.py b/lmdeploy/pytorch/kernels/cuda/moe/ep.py index 53ee8f5a19..cde008b8ee 100644 --- a/lmdeploy/pytorch/kernels/cuda/moe/ep.py +++ b/lmdeploy/pytorch/kernels/cuda/moe/ep.py @@ -142,20 +142,22 @@ def _fwd_kernel_ep_gather( output_tensor_stride1, topk_num: tl.constexpr, BLOCK_D: tl.constexpr, + FP32_ACC: tl.constexpr, ): cur_block = tl.program_id(0) start_cur_token = tl.program_id(1) grid_num = tl.num_programs(1) + compute_dtype = tl.float32 if FP32_ACC else output_tensor.dtype.element_ty for cur_token in range(start_cur_token, total_token_num, grid_num): off_d = tl.arange(0, BLOCK_D) - accumulator = tl.zeros([BLOCK_D], dtype=tl.float32) + accumulator = tl.zeros([BLOCK_D], dtype=compute_dtype) for topk_index in range(0, topk_num): expert_id = tl.load(recv_topk_ids + cur_token * recv_topk_ids_stride0 + topk_index) if expert_id >= 0: source_token_index = tl.load(input_index + cur_token * input_index_stride0 + topk_index) acc_weight = tl.load(recv_topk_weight + cur_token * recv_topk_weight_stride0 + topk_index) tmp = tl.load(input_tensor + source_token_index * input_tensor_stride0 + cur_block * BLOCK_D + off_d) - accumulator += tmp.to(tl.float32) * acc_weight.to(tl.float32) + accumulator += tmp.to(compute_dtype) * acc_weight.to(compute_dtype) tl.store( output_tensor + cur_token * output_tensor_stride0 + cur_block * BLOCK_D + off_d, accumulator.to(output_tensor.dtype.element_ty), @@ -169,6 +171,7 @@ def ep_gather( recv_topk_weight: torch.Tensor, input_index: torch.Tensor, output_tensor: torch.Tensor, + fp32_acc: bool = False, ): BLOCK_D = 1024 # block size of quantization num_warps = 2 @@ -196,6 +199,7 @@ def ep_gather( topk_num=recv_topk_ids.shape[1], num_warps=num_warps, BLOCK_D=BLOCK_D, + FP32_ACC=fp32_acc, ) return diff --git a/lmdeploy/pytorch/kernels/cuda/moe/ep_fp8.py b/lmdeploy/pytorch/kernels/cuda/moe/ep_fp8.py index a092381d29..6fd9c64723 100644 --- a/lmdeploy/pytorch/kernels/cuda/moe/ep_fp8.py +++ b/lmdeploy/pytorch/kernels/cuda/moe/ep_fp8.py @@ -156,6 +156,7 @@ def fused_moe_v3_fp8( num_recv_tokens_per_expert: list[int] | None, act_func: Callable | None = None, scale_fmt: str | None = None, + fp32_acc: bool = False, ): hidden_states_fp8, hidden_states_scale = hidden_states_fp8 if num_recv_tokens_per_expert is None: @@ -167,7 +168,7 @@ def fused_moe_v3_fp8( m, k = hidden_states_fp8.size() n = w13_weight_fp8[0].size(1) block_size = k // hidden_states_scale.size(1) - # ep_gather accumulates in FP32 and casts once when storing BF16 for DeepEP. + # Keep DeepEP's BF16 transport dtype independent of local accumulation precision. gather_out = torch.empty_like(hidden_states_fp8, dtype=torch.bfloat16) input_tensor = torch.empty((all_tokens, k), device=hidden_states_fp8.device, dtype=hidden_states_fp8.dtype) input_tensor_scale = torch.empty((all_tokens, k // block_size), @@ -197,6 +198,5 @@ def fused_moe_v3_fp8( down_input_scale = get_mn_major_tma_aligned_tensor(down_input_scale) down_output = torch.empty((all_tokens, k), device=gather_out.device, dtype=torch.bfloat16) _deepgemm_grouped_fp8_nt_contiguous((down_input_fp8, down_input_scale), w2_weight_fp8, down_output, m_indices) - ep_gather(down_output, topk_idx, topk_weights, output_index, gather_out) - # DeepEP transports BF16 partial sums; only the local reduction is FP32. + ep_gather(down_output, topk_idx, topk_weights, output_index, gather_out, fp32_acc=fp32_acc) return gather_out From 08d7a74f08e421827ca370c9d9305df6c758e353 Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Mon, 28 Sep 2026 08:34:50 +0000 Subject: [PATCH 30/39] refactor(pytorch): remove unused sparse MLA kernel and padding hook --- .../kernels/cuda/sparse_mla_tilelang.py | 233 ------------------ lmdeploy/pytorch/models/deepseek_v32.py | 4 +- 2 files changed, 1 insertion(+), 236 deletions(-) delete mode 100644 lmdeploy/pytorch/kernels/cuda/sparse_mla_tilelang.py diff --git a/lmdeploy/pytorch/kernels/cuda/sparse_mla_tilelang.py b/lmdeploy/pytorch/kernels/cuda/sparse_mla_tilelang.py deleted file mode 100644 index 247d8876cd..0000000000 --- a/lmdeploy/pytorch/kernels/cuda/sparse_mla_tilelang.py +++ /dev/null @@ -1,233 +0,0 @@ -# Copyright (c) OpenMMLab. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 -"""TileLang sparse MLA decode attention for CUDA BF16 tensors. - -The kernel schedule in this file is adapted from SGLang's sparse_attention_fwd_kernel_v1 at commit 9e692c9216c3, -distributed under the Apache License 2.0: - -https://github.com/sgl-project/sglang/blob/9e692c9216c3/python/sglang/kernels/ops/attention/dsa/tilelang_kernel.py - -Only the CUDA BF16 path used by GLM-5.3-Flash is retained. The public wrapper -uses generic sparse-MLA names and has no SGLang runtime dependency. -""" - -import tilelang -import tilelang.language as T -import torch - -tilelang.set_log_level('WARNING') - - -@tilelang.jit( - out_idx=[-1], - pass_configs={ - tilelang.PassConfigKey.TL_DISABLE_TMA_LOWER: True, - tilelang.PassConfigKey.TL_DISABLE_WARP_SPECIALIZED: True, - }, -) -def _sparse_mla_bf16_fwd_kernel( - num_heads, - dim, - topk, - *, - storage_dim=None, - kv_group=1, - sm_scale=None, - is_causal=True, - block_I=64, - num_stages=2, - threads=256, -): - if storage_dim is None: - storage_dim = dim - assert storage_dim >= dim - assert ( - dim == tilelang.math.next_power_of_2(dim) or dim % 64 == 0 - ), f"dim={dim} must be a power of 2 or a multiple of 64" - assert is_causal, 'non-causal is not supported' - assert ( - topk % block_I == 0 - ), 'otherwise will load some index=0 thus causing wrong kv to be loaded' - if sm_scale is None: - sm_scale = (1.0 / dim) ** 0.5 * 1.44269504 # log2(e) - else: - sm_scale = sm_scale * 1.44269504 # log2(e) - - batch = T.symbolic('batch') - seq_len = T.symbolic('seq_len') - seq_len_kv = T.symbolic('seq_len_kv') - - head_kv = num_heads // kv_group - q_shape = [batch, seq_len, num_heads, dim] - kv_shape = [batch, seq_len_kv, kv_group, storage_dim] - o_shape = [batch, seq_len, num_heads, dim] - indices_shape = [batch, seq_len, kv_group, topk] - indices_dtype = 'int32' - dtype = 'bfloat16' - accum_dtype = 'float' - - H = head_kv - padded_H = max(tilelang.math.next_power_of_2(head_kv), 16) - if padded_H != H: - assert kv_group == 1 - BI = block_I - NI = tilelang.cdiv(topk, block_I) - D = dim - - if head_kv > 64: - assert head_kv % 64 == 0, 'head_kv should be a multiple of 64' - REPLICATE_H = head_kv // 64 - else: - REPLICATE_H = 1 - - H_per_block = padded_H if REPLICATE_H == 1 else 64 - - @T.prim_func - def main( - Q: T.Tensor(q_shape, dtype), # type: ignore - KV: T.Tensor(kv_shape, dtype), # type: ignore - Indices: T.Tensor(indices_shape, indices_dtype), # type: ignore - Output: T.Tensor(o_shape, dtype), # type: ignore - ): - with T.Kernel(seq_len * REPLICATE_H, batch, kv_group, threads=threads) as ( - bx, - by, - bz, - ): - Q_shared = T.alloc_shared([H_per_block, D], dtype) - KV_shared = T.alloc_shared([BI, D], dtype) - O_shared = T.alloc_shared([H_per_block, D], dtype) - mask = T.alloc_fragment([BI], 'bool') - - acc_o = T.alloc_fragment([H_per_block, D], accum_dtype) - acc_s = T.alloc_fragment([H_per_block, BI], accum_dtype) - S_shared = T.alloc_shared([H_per_block, BI], dtype) - sumexp = T.alloc_fragment([H_per_block], accum_dtype) - sumexp_i = T.alloc_fragment([H_per_block], accum_dtype) - alpha = T.alloc_fragment([H_per_block], accum_dtype) - m_i = T.alloc_fragment([H_per_block], accum_dtype) - m_i_prev = T.alloc_fragment([H_per_block], accum_dtype) - - T.fill(acc_o, 0) - T.fill(sumexp, 0) - T.fill(m_i, -(2**30)) # avoid -inf - inf to cause nan - - b_i, g_i = by, bz - s_i = bx if REPLICATE_H == 1 else (bx // REPLICATE_H) - - H0 = g_i * padded_H + (0 if REPLICATE_H == 1 else (bx % REPLICATE_H) * 64) - H1 = H0 + H_per_block - - T.copy(Q[b_i, s_i, H0:H1, :D], Q_shared) - - for i_i in T.Pipelined(NI, num_stages=num_stages): - - for bi_i in T.Parallel(BI): - mask[bi_i] = Indices[b_i, s_i, g_i, i_i * BI + bi_i] >= 0 - - for bi_i, d_i in T.Parallel(BI, D): - KV_shared[bi_i, d_i] = KV[ - b_i, Indices[b_i, s_i, g_i, i_i * BI + bi_i], g_i, d_i - ] - - for h_i, bi_i in T.Parallel(H_per_block, BI): - acc_s[h_i, bi_i] = T.if_then_else( - mask[bi_i], 0, -T.infinity(acc_s.dtype) - ) - T.gemm( - Q_shared, - KV_shared, - acc_s, - transpose_B=True, - policy=T.GemmWarpPolicy.FullCol, - ) - T.copy(m_i, m_i_prev) - T.reduce_max(acc_s, m_i, dim=1, clear=False) - for h_i in T.Parallel(H_per_block): - alpha[h_i] = T.exp2((m_i_prev[h_i] - m_i[h_i]) * sm_scale) - for h_i, bi_i in T.Parallel(H_per_block, BI): - acc_s[h_i, bi_i] = T.exp2( - acc_s[h_i, bi_i] * sm_scale - m_i[h_i] * sm_scale - ) - T.reduce_sum(acc_s, sumexp_i, dim=1) # is this a accumulate operator? - for h_i in T.Parallel(H_per_block): - sumexp[h_i] = sumexp[h_i] * alpha[h_i] + sumexp_i[h_i] - for h_i, d_i in T.Parallel(H_per_block, D): - acc_o[h_i, d_i] = acc_o[h_i, d_i] * alpha[h_i] - - T.copy(acc_s, S_shared) - T.gemm(S_shared, KV_shared, acc_o, policy=T.GemmWarpPolicy.FullCol) - - # Rescale - for h_i, d_i in T.Parallel(H_per_block, D): - acc_o[h_i, d_i] /= sumexp[h_i] - for h_i in T.Parallel(H_per_block): - sumexp[h_i] = T.log2(sumexp[h_i]) + m_i[h_i] * sm_scale - - T.copy(acc_o, O_shared) - T.copy(acc_o, Output[b_i, s_i, H0:H1, :]) - - return main - - -def sparse_mla_bf16_fwd(q: torch.Tensor, - kv: torch.Tensor, - indices: torch.Tensor, - sm_scale: float) -> torch.Tensor: - """Run sparse MLA decode attention on CUDA BF16 tensors. - - Args: - q: Query tensor with shape [tokens, local_heads, 512]. - kv: Paged latent KV tensor with shape [slots, 1, storage_dim]. The - supported storage widths are 512 and 576; only the first 512 - elements participate in attention. - indices: Physical slot indices with shape [tokens, 1, topk]. - Invalid entries must be -1 and topk must be padded to a - multiple of 64. - sm_scale: Attention scale before the kernel's base-2 conversion. - - Returns: - A BF16 tensor with shape [tokens, local_heads, 512]. - """ - if torch.version.hip is not None: - raise RuntimeError('sparse_mla_bf16_fwd only supports CUDA') - if not (q.is_cuda and kv.is_cuda and indices.is_cuda): - raise ValueError('q, kv, and indices must be CUDA tensors') - if q.dtype != torch.bfloat16 or kv.dtype != torch.bfloat16: - raise TypeError('q and kv must have dtype torch.bfloat16') - if indices.dtype != torch.int32: - raise TypeError('indices must have dtype torch.int32') - if q.ndim != 3 or q.shape[1] <= 0 or q.shape[2] != 512: - raise ValueError( - 'q must have shape [tokens, positive_local_heads, 512], ' - f'got {tuple(q.shape)}') - if kv.ndim != 3 or kv.shape[1] != 1 or kv.shape[2] not in (512, 576): - raise ValueError( - 'kv must have shape [slots, 1, storage_dim] with storage_dim ' - f'in (512, 576), got {tuple(kv.shape)}') - if (indices.ndim != 3 or indices.shape[0] != q.shape[0] - or indices.shape[1] != 1): - raise ValueError( - 'indices must have shape [tokens, 1, topk] matching q, ' - f'got {tuple(indices.shape)}') - if q.device != kv.device or q.device != indices.device: - raise ValueError('q, kv, and indices must be on the same CUDA device') - - topk = indices.shape[-1] - if topk == 0 or topk % 64 != 0: - raise ValueError(f'topk must be a positive multiple of 64, got {topk}') - - if not kv.is_contiguous(): - raise ValueError( - 'kv must be a zero-copy contiguous view of the complete paged cache') - q = q.contiguous() - indices = indices.contiguous() - kernel = _sparse_mla_bf16_fwd_kernel( - num_heads=q.shape[1], - dim=512, - topk=topk, - storage_dim=kv.shape[-1], - sm_scale=sm_scale, - ) - output = kernel(q.unsqueeze(0), kv.unsqueeze(0), indices.unsqueeze(0)) - return output.squeeze(0) diff --git a/lmdeploy/pytorch/models/deepseek_v32.py b/lmdeploy/pytorch/models/deepseek_v32.py index a7431ac6be..c8dfb15fe9 100644 --- a/lmdeploy/pytorch/models/deepseek_v32.py +++ b/lmdeploy/pytorch/models/deepseek_v32.py @@ -251,7 +251,6 @@ def forward(self, class DeepseekV32Attention(DeepseekV2Attention): use_sparse_mla = True - mla_head_padding = 0 def __init__(self, config: Any, @@ -358,8 +357,7 @@ def __init__(self, self.softmax_scale = self.softmax_scale * mscale * mscale self.attn_fwd = Attention(self.num_heads, - config.kv_lora_rank + self.qk_rope_head_dim - + type(self).mla_head_padding, + config.kv_lora_rank + self.qk_rope_head_dim, scale=self.softmax_scale, num_kv_heads=num_key_value_heads, v_head_size=config.kv_lora_rank, From 18a8c917120e60315a7b628de5b1b8331d2bc755 Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Mon, 28 Sep 2026 09:19:08 +0000 Subject: [PATCH 31/39] fix(pytorch): preserve MoE reduction defaults with explicit GLM FP32 opt-in Restore moe_reduce's positional/keyword fp32_acc=False argument and the legacy weighted-product precision for default callers. Keep output_scale keyword-only and apply it after expert reduction before the output cast. GLM explicitly sets fused_moe_fp32_acc=True. Propagate it through the shared model, builder, typed specs, BF16/blocked-FP8 backends and kernels. DeepEP Normal uses this flag independently of routing scale. Reject unsupported non-default requests instead of silently dropping them. Restore compressed_tensors_w4a16.py exactly to main: its existing explicit FP32 calls are compatible again. The PR now changes 65 files, with no new production or test files introduced by this update. Validation: 22 focused H200/CPU checks passed, including exact legacy and previous GLM reduction parity, positional/keyword API compatibility, FP32 scaling, CUDA graph replay, parameter propagation, unsupported backend rejection and actual BF16/blocked-FP8 expert pipelines. Ruff and git diff --check passed. Full-model and distributed inference not rerun. --- .../pytorch/backends/cuda/moe/blocked_fp8.py | 16 +++++++++++----- lmdeploy/pytorch/backends/cuda/moe/default.py | 16 +++++++++++----- lmdeploy/pytorch/backends/dlinfer/moe.py | 4 ++-- lmdeploy/pytorch/backends/moe.py | 2 ++ .../kernels/cuda/compressed_tensors_w4a16.py | 4 ++-- .../pytorch/kernels/cuda/moe/blocked_fp8.py | 4 +++- lmdeploy/pytorch/kernels/cuda/moe/fused_moe.py | 18 +++++++++++++----- lmdeploy/pytorch/models/deepseek_v2.py | 2 ++ lmdeploy/pytorch/models/glm5_next.py | 1 + lmdeploy/pytorch/nn/moe/__init__.py | 12 ++++++++---- lmdeploy/pytorch/nn/moe/blocked_fp8.py | 4 +++- lmdeploy/pytorch/nn/moe/default.py | 4 +++- 12 files changed, 61 insertions(+), 26 deletions(-) diff --git a/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py b/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py index bba12e1821..8f42e954a6 100644 --- a/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py +++ b/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py @@ -311,7 +311,8 @@ def __init__(self, renormalize: bool = False, block_size: int = 128, out_dtype: torch.dtype = torch.float16, - output_scale: float = 1.0): + output_scale: float = 1.0, + fp32_acc: bool = False): super().__init__() self.num_experts = num_experts self.top_k = top_k @@ -319,6 +320,7 @@ def __init__(self, self.block_size = block_size self.out_dtype = out_dtype self.output_scale = output_scale + self.fp32_acc = fp32_acc def ep_expert_list(self, world_size: int, rank: int): """Experts list of current rank.""" @@ -368,7 +370,8 @@ def forward(self, num_experts=num_experts, renormalize=self.renormalize, act_func=act_func, - output_scale=self.output_scale) + output_scale=self.output_scale, + fp32_acc=self.fp32_acc) output = output.unflatten(0, input_size[:-1]) return output @@ -389,8 +392,9 @@ def __init__(self, fp8_dtype: torch.dtype = torch.float8_e4m3fn, num_max_dispatch_tokens_per_rank: int = 128, layer_idx: int = 0, - output_scale: float = 1.0): - super().__init__(top_k, num_experts, renormalize, block_size, out_dtype, output_scale) + output_scale: float = 1.0, + fp32_acc: bool = False): + super().__init__(top_k, num_experts, renormalize, block_size, out_dtype, output_scale, fp32_acc) self.num_experts = num_experts self.ep_size = ep_size self.ep_group = ep_group @@ -563,7 +567,7 @@ def fusedmoe_build(self, low_latency_mode: bool = False): scale_fmt=self.scale_fmt, layer_idx=self.layer_idx, num_max_dispatch_tokens_per_rank=self.num_max_dispatch_tokens_per_rank, - fp32_acc=self.output_scale != 1.0, + fp32_acc=self.fp32_acc, chunk_size=16 * 1024) return deepep_moe @@ -584,6 +588,7 @@ def _build_fused_moe_blocked_f8(spec: FusedMoEBlockedF8BuildSpec) -> FusedMoEBlo num_max_dispatch_tokens_per_rank=spec.num_max_dispatch_tokens_per_rank, layer_idx=spec.layer_idx, output_scale=spec.output_scale, + fp32_acc=spec.fp32_acc, ) else: impl = TritonFusedMoEBlockedF8Impl( @@ -593,6 +598,7 @@ def _build_fused_moe_blocked_f8(spec: FusedMoEBlockedF8BuildSpec) -> FusedMoEBlo block_size=spec.block_size, out_dtype=spec.output_dtype, output_scale=spec.output_scale, + fp32_acc=spec.fp32_acc, ) impl.set_scale_fmt(spec.scale_fmt) return impl diff --git a/lmdeploy/pytorch/backends/cuda/moe/default.py b/lmdeploy/pytorch/backends/cuda/moe/default.py index fe60d1c859..434f8a898e 100644 --- a/lmdeploy/pytorch/backends/cuda/moe/default.py +++ b/lmdeploy/pytorch/backends/cuda/moe/default.py @@ -26,11 +26,13 @@ def __init__(self, top_k: int, num_experts: int, renormalize: bool = False, - output_scale: float = 1.0): + output_scale: float = 1.0, + fp32_acc: bool = False): self.num_experts = num_experts self.top_k = top_k self.renormalize = renormalize self.output_scale = output_scale + self.fp32_acc = fp32_acc def update_weights(self, gate_up_weights: torch.Tensor, down_weights: torch.Tensor): gate_up_weights = gate_up_weights.transpose(1, 2).contiguous().transpose(1, 2) @@ -73,7 +75,8 @@ def forward(self, num_experts=num_experts, renormalize=self.renormalize, act_func=act_func, - output_scale=self.output_scale) + output_scale=self.output_scale, + fp32_acc=self.fp32_acc) # modify from dlblas: https://github.com/DeepLink-org/DLBlas @@ -385,10 +388,11 @@ def __init__( layer_idx: int = 0, out_dtype: torch.dtype = torch.bfloat16, num_max_dispatch_tokens_per_rank: int = 128, + fp32_acc: bool = False, ): - super().__init__(top_k, num_experts, renormalize, output_scale) - if output_scale != 1.0: - raise NotImplementedError('DeepEP MoE does not support output_scale.') + super().__init__(top_k, num_experts, renormalize, output_scale, fp32_acc) + if fp32_acc or output_scale != 1.0: + raise NotImplementedError('DeepEP BF16 MoE does not support fp32_acc or output_scale.') self.num_experts = num_experts self.ep_size = ep_size self.ep_group = ep_group @@ -550,6 +554,7 @@ def _build_fused_moe(spec: FusedMoEBuildSpec) -> FusedMoEImpl: hidden_dim=spec.hidden_dim, renormalize=spec.renormalize, output_scale=spec.output_scale, + fp32_acc=spec.fp32_acc, layer_idx=spec.layer_idx, out_dtype=spec.output_dtype, num_max_dispatch_tokens_per_rank=spec.num_max_dispatch_tokens_per_rank, @@ -559,4 +564,5 @@ def _build_fused_moe(spec: FusedMoEBuildSpec) -> FusedMoEImpl: num_experts=spec.num_experts, renormalize=spec.renormalize, output_scale=spec.output_scale, + fp32_acc=spec.fp32_acc, ) diff --git a/lmdeploy/pytorch/backends/dlinfer/moe.py b/lmdeploy/pytorch/backends/dlinfer/moe.py index c7f4ebbeca..0cbd34a546 100644 --- a/lmdeploy/pytorch/backends/dlinfer/moe.py +++ b/lmdeploy/pytorch/backends/dlinfer/moe.py @@ -117,9 +117,9 @@ def forward(self, def _build_fused_moe(spec: FusedMoEBuildSpec) -> FusedMoEImpl: """Build a DLINFER fused MoE implementation.""" - if spec.output_scale != 1.0: + if spec.fp32_acc or spec.output_scale != 1.0: raise NotImplementedError( - 'DLINFER fused MoE does not support output_scale.') + 'DLINFER fused MoE does not support fp32_acc or output_scale.') return DlinferFusedMoEImpl( top_k=spec.top_k, num_experts=spec.num_experts, diff --git a/lmdeploy/pytorch/backends/moe.py b/lmdeploy/pytorch/backends/moe.py index 71ce73812f..810af52648 100644 --- a/lmdeploy/pytorch/backends/moe.py +++ b/lmdeploy/pytorch/backends/moe.py @@ -74,6 +74,7 @@ class FusedMoEBuildSpec(BuildSpec[FusedMoEImpl]): output_dtype: torch.dtype num_max_dispatch_tokens_per_rank: int output_scale: float = 1.0 + fp32_acc: bool = False class FusedMoEW8A8Impl(ABC): @@ -258,6 +259,7 @@ class FusedMoEBlockedF8BuildSpec(BuildSpec[FusedMoEBlockedF8Impl]): custom_gateup_act: bool scale_fmt: str | None output_scale: float = 1.0 + fp32_acc: bool = False class FusedMoEV4FP4Impl(ABC): diff --git a/lmdeploy/pytorch/kernels/cuda/compressed_tensors_w4a16.py b/lmdeploy/pytorch/kernels/cuda/compressed_tensors_w4a16.py index 18d1279379..b84c9b25b8 100644 --- a/lmdeploy/pytorch/kernels/cuda/compressed_tensors_w4a16.py +++ b/lmdeploy/pytorch/kernels/cuda/compressed_tensors_w4a16.py @@ -915,7 +915,7 @@ def fused_moe_w4a16( num_bits=num_bits, group_size=group_size, ) - return moe_reduce(expert_output, topk_weights) + return moe_reduce(expert_output, topk_weights, fp32_acc=True) # PyTorch promotes int32 cumsums to int64. Normalize only paths that use # the shared sorter so all routing metadata keeps a homogeneous dtype. @@ -1012,4 +1012,4 @@ def fused_moe_w4a16( ) if valid_routes is not None: expert_output.masked_fill_(~valid_routes[..., None], 0) - return moe_reduce(expert_output, topk_weights) + return moe_reduce(expert_output, topk_weights, fp32_acc=True) diff --git a/lmdeploy/pytorch/kernels/cuda/moe/blocked_fp8.py b/lmdeploy/pytorch/kernels/cuda/moe/blocked_fp8.py index c5fe21a9a8..641ecb8f1f 100644 --- a/lmdeploy/pytorch/kernels/cuda/moe/blocked_fp8.py +++ b/lmdeploy/pytorch/kernels/cuda/moe/blocked_fp8.py @@ -697,7 +697,8 @@ def fused_moe_blocked_fp8(input: torch.Tensor, renormalize: bool = False, act_func: Callable = None, *, - output_scale: float = 1.0) -> torch.Tensor: + output_scale: float = 1.0, + fp32_acc: bool = False) -> torch.Tensor: """Fused moe.""" device = input.device M = input.size(0) @@ -824,5 +825,6 @@ def fused_moe_blocked_fp8(input: torch.Tensor, ) ret = moe_reduce(intermediate_cache2, topk_weights, + fp32_acc=fp32_acc, output_scale=output_scale) return ret diff --git a/lmdeploy/pytorch/kernels/cuda/moe/fused_moe.py b/lmdeploy/pytorch/kernels/cuda/moe/fused_moe.py index 44b548443a..75baa38734 100644 --- a/lmdeploy/pytorch/kernels/cuda/moe/fused_moe.py +++ b/lmdeploy/pytorch/kernels/cuda/moe/fused_moe.py @@ -942,6 +942,7 @@ def _moe_reduce_kernel( stride_wk: tl.constexpr, stride_om, stride_on: tl.constexpr, + fp32_acc: tl.constexpr, K: tl.constexpr, N: tl.constexpr, BLOCK_K: tl.constexpr, @@ -966,8 +967,11 @@ def _moe_reduce_kernel( h = tl.load(h_ptrs, mask=mask_h, other=0.0) w = tl.load(weights_ptrs, mask=mask_k, other=0.0) - h = h.to(tl.float32) - w = w.to(tl.float32) + if fp32_acc: + h = h.to(tl.float32) + w = w.to(tl.float32) + else: + w = w.to(h.dtype) wh = h * w[:, None] o = wh.sum(axis=0) @@ -976,8 +980,9 @@ def _moe_reduce_kernel( tl.store(o_ptrs, o, mask=mask_n) -def moe_reduce(hidden_states: torch.Tensor, topk_weights: torch.Tensor, *, output_scale: float = 1.0) -> torch.Tensor: - """Weight and reduce experts, optionally scaling before the output cast.""" +def moe_reduce(hidden_states: torch.Tensor, topk_weights: torch.Tensor, fp32_acc: bool = False, + *, output_scale: float = 1.0) -> torch.Tensor: + """Weight and reduce experts with optional FP32 products and output scaling.""" assert hidden_states.dim() == 3 assert topk_weights.dim() == 2 assert hidden_states.size(0) == topk_weights.size(0) @@ -1002,6 +1007,7 @@ def moe_reduce(hidden_states: torch.Tensor, topk_weights: torch.Tensor, *, outpu topk_weights.stride(1), out.stride(0), out.stride(1), + fp32_acc, K, N, BLOCK_K, @@ -1025,7 +1031,8 @@ def fused_moe(hidden_states: torch.Tensor, num_experts: int = None, renormalize: bool = False, act_func: Callable = None, - output_scale: float = 1.0) -> torch.Tensor: + output_scale: float = 1.0, + fp32_acc: bool = False) -> torch.Tensor: """Fused moe.""" M = hidden_states.size(0) E, N, _ = w1.shape @@ -1135,5 +1142,6 @@ def fused_moe(hidden_states: torch.Tensor, ret = moe_reduce(intermediate_cache2, topk_weights, + fp32_acc=fp32_acc, output_scale=output_scale) return ret diff --git a/lmdeploy/pytorch/models/deepseek_v2.py b/lmdeploy/pytorch/models/deepseek_v2.py index 3ebc914112..f4727231f3 100644 --- a/lmdeploy/pytorch/models/deepseek_v2.py +++ b/lmdeploy/pytorch/models/deepseek_v2.py @@ -709,6 +709,7 @@ class DeepseekV2MoE(nn.Module): fused_moe_act_func = None fused_moe_output_scale = 1.0 + fused_moe_fp32_acc = False router_routed_scaling_factor = None shared_expert_cls = None @@ -772,6 +773,7 @@ def __init__(self, layer_idx=layer_idx, act_func=type(self).fused_moe_act_func, output_scale=type(self).fused_moe_output_scale, + fp32_acc=type(self).fused_moe_fp32_acc, prefix=add_prefix('experts', prefix), ) self.shared_experts = None diff --git a/lmdeploy/pytorch/models/glm5_next.py b/lmdeploy/pytorch/models/glm5_next.py index a9f6b76ebb..d690499574 100644 --- a/lmdeploy/pytorch/models/glm5_next.py +++ b/lmdeploy/pytorch/models/glm5_next.py @@ -511,6 +511,7 @@ class Glm5NextMoE(DeepseekV2MoE): # scales FP32 combine weights; shared experts remain unscaled. router_routed_scaling_factor = 1.0 fused_moe_output_scale = 2.5 + fused_moe_fp32_acc = True shared_expert_cls = Glm5NextMLP def __init__(self, config: Any, layer_idx: int, *args, **kwargs): diff --git a/lmdeploy/pytorch/nn/moe/__init__.py b/lmdeploy/pytorch/nn/moe/__init__.py index b38678526f..7214dabece 100644 --- a/lmdeploy/pytorch/nn/moe/__init__.py +++ b/lmdeploy/pytorch/nn/moe/__init__.py @@ -26,6 +26,7 @@ def build_fused_moe( prefix: str = '', *, output_scale: float = 1.0, + fp32_acc: bool = False, ): """Fused moe builder.""" quant_method = None @@ -48,11 +49,12 @@ def build_fused_moe( layer_idx=layer_idx, act_func=act_func, output_scale=output_scale, + fp32_acc=fp32_acc, ) if quant_method == 'smooth_quant': - if output_scale != 1.0: - raise NotImplementedError('W8A8 MoE does not support output_scale.') + if fp32_acc or output_scale != 1.0: + raise NotImplementedError('W8A8 MoE does not support fp32_acc or output_scale.') assert not bias, 'Quant model does not support bias for now.' assert act_func is None, ('Quant model does not support activation function for now.') from .w8a8 import FusedMoEW8A8 @@ -74,8 +76,8 @@ def build_fused_moe( ) if is_static_per_tensor: - if output_scale != 1.0: - raise NotImplementedError('Static FP8 MoE does not support output_scale.') + if fp32_acc or output_scale != 1.0: + raise NotImplementedError('Static FP8 MoE does not support fp32_acc or output_scale.') assert not bias, ( 'Static FP8 MoE does not support bias.' ) @@ -115,8 +117,10 @@ def build_fused_moe( layer_idx=layer_idx, act_func=act_func, output_scale=output_scale, + fp32_acc=fp32_acc, ) elif quant_method == 'compressed-tensors': + # W4A16 already requests FP32 expert reduction for every invocation. if output_scale != 1.0: raise NotImplementedError('W4A16 MoE does not support output_scale.') if bias: diff --git a/lmdeploy/pytorch/nn/moe/blocked_fp8.py b/lmdeploy/pytorch/nn/moe/blocked_fp8.py index 157d9273a4..da292b0513 100644 --- a/lmdeploy/pytorch/nn/moe/blocked_fp8.py +++ b/lmdeploy/pytorch/nn/moe/blocked_fp8.py @@ -160,7 +160,8 @@ def __init__(self, all_reduce: bool = True, layer_idx: int = 0, act_func: Callable = None, - output_scale: float = 1.0): + output_scale: float = 1.0, + fp32_acc: bool = False): device = device or torch.device('cpu') dtype = dtype or torch.float16 @@ -194,6 +195,7 @@ def __init__(self, custom_gateup_act=act_func is not None, scale_fmt=scale_fmt, output_scale=output_scale, + fp32_acc=fp32_acc, ), enable_deterministic=build_ctx.enable_deterministic, ) diff --git a/lmdeploy/pytorch/nn/moe/default.py b/lmdeploy/pytorch/nn/moe/default.py index a54d42e233..02d7da928d 100644 --- a/lmdeploy/pytorch/nn/moe/default.py +++ b/lmdeploy/pytorch/nn/moe/default.py @@ -131,7 +131,8 @@ def __init__(self, all_reduce: bool = True, layer_idx: int = 0, act_func: Callable = None, - output_scale: float = 1.0): + output_scale: float = 1.0, + fp32_acc: bool = False): device = device or torch.device('cpu') dtype = dtype or torch.float16 @@ -160,6 +161,7 @@ def __init__(self, output_dtype=torch.bfloat16, num_max_dispatch_tokens_per_rank=build_ctx.deep_ep_max_tokens_per_rank, output_scale=output_scale, + fp32_acc=fp32_acc, ), enable_deterministic=build_ctx.enable_deterministic, ) From 200a559150e0b09c12d14faca2ce36beda48a5a0 Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Mon, 28 Sep 2026 13:49:53 +0000 Subject: [PATCH 32/39] fix(pytorch): bound and reserve GLM KPool prefill score memory --- lmdeploy/pytorch/backends/cuda/kpool.py | 32 ++++++++++++++------ lmdeploy/pytorch/backends/cuda/nsa.py | 16 +++++++++- lmdeploy/pytorch/config.py | 2 ++ lmdeploy/pytorch/configurations/glm5_next.py | 1 + lmdeploy/pytorch/engine/executor/base.py | 3 +- 5 files changed, 43 insertions(+), 11 deletions(-) diff --git a/lmdeploy/pytorch/backends/cuda/kpool.py b/lmdeploy/pytorch/backends/cuda/kpool.py index 87067dd78a..872a99b408 100644 --- a/lmdeploy/pytorch/backends/cuda/kpool.py +++ b/lmdeploy/pytorch/backends/cuda/kpool.py @@ -8,6 +8,7 @@ import torch from torch import Tensor +from lmdeploy.pytorch import envs as _envs from lmdeploy.pytorch.kernels.cuda.fill_kv_cache import fill_indexed_key_cache from lmdeploy.pytorch.kernels.cuda.flatten_kv_cache import flatten_kv_cache from lmdeploy.pytorch.kernels.cuda.kpool import ( @@ -29,6 +30,7 @@ ) from .gated_delta_rule import _state_scatter +from .nsa import _get_max_score_rows # Fuse the existing integer index expansion instead of materializing its # [tokens, topk] masks and int64 temporaries separately during prefill or decode. @@ -138,15 +140,27 @@ def kpool_select_prefill_cuda(query_fp8, query_weight, packed_cache, keys.view(torch.uint8).unsqueeze(2), scales.view(torch.uint8).unsqueeze(2), counts, blocks, start_loc=starts, out_size=max(1, kv_flatten_size // pool_size)) max_groups = block_offsets.size(1) * keys.size(1) // pool_size - logits = _get_deep_gemm().fp8_mqa_logits( - query_fp8.contiguous(), - (flat_keys[0].view(torch.float8_e4m3fn), flat_scales[0].view(torch.float32).flatten()), - query_weight.contiguous(), query_starts, query_ends, - clean_logits=False, max_seqlen_k=max_groups) - # Compressed logits are request-local. The selector masks the unwritten - # suffix using each query's causal length, including zero-length rows. - selected = kpool_select_groups_cuda( - logits, lengths, group_topk=topk // pool_size, max_group_length=max_groups) + max_rows = _get_max_score_rows( + max_groups, _envs.dsa_indexer_max_logits_mb * (1 << 20), num_heads=query_fp8.size(1)) + flat_kv = (flat_keys[0].view(torch.float8_e4m3fn), flat_scales[0].view(torch.float32).flatten()) + selected = None + if rows > max_rows: + selected = torch.empty((rows, topk // pool_size), device=query_fp8.device, dtype=torch.int32) + for start in range(0, rows, max_rows): + row_slice = slice(start, min(start + max_rows, rows)) + logits = _get_deep_gemm().fp8_mqa_logits( + query_fp8[row_slice].contiguous(), flat_kv, query_weight[row_slice].contiguous(), + query_starts[row_slice], query_ends[row_slice], + clean_logits=False, max_seqlen_k=max_groups) + # Keep request-local causal lengths and deterministic Top-K unchanged. + chunk = kpool_select_groups_cuda( + logits, lengths[row_slice], group_topk=topk // pool_size, max_group_length=max_groups) + if selected is None: + selected = chunk + else: + selected[row_slice].copy_(chunk) + # Release before allocating the next score chunk, including its padding. + del logits, chunk return kpool_expand_groups_cuda(selected, lengths, pool_size, topk, seq_lens=seq) diff --git a/lmdeploy/pytorch/backends/cuda/nsa.py b/lmdeploy/pytorch/backends/cuda/nsa.py index 976d884c41..3312c8b1d4 100644 --- a/lmdeploy/pytorch/backends/cuda/nsa.py +++ b/lmdeploy/pytorch/backends/cuda/nsa.py @@ -42,7 +42,7 @@ logger = get_logger('lmdeploy') -def _get_max_score_rows(max_kv_seqlen: int, max_logits_bytes: int) -> int: +def _get_max_score_rows(max_kv_seqlen: int, max_logits_bytes: int, *, num_heads: int | None = None) -> int: """Return the query rows fitting in a bounded FP32 score tensor.""" if max_kv_seqlen <= 0: return 1 @@ -52,6 +52,20 @@ def _get_max_score_rows(max_kv_seqlen: int, max_logits_bytes: int) -> int: # Bounding flattened KV alone therefore does not bound the M * N logits # allocation; limit M so its FP32 payload stays within the runtime budget. _fp32_bytes = 4 + if num_heads is not None: + # Compressed DeepGEMM logits allocate ceil(M / block_q) * block_q + # rows with a 256-element (1024-byte) aligned FP32 row stride. + if not 0 < num_heads <= 128 or num_heads % 4: + raise ValueError('DeepGEMM score heads must be a multiple of four in [4, 128].') + block_q = 128 // num_heads + row_bytes = ((max_kv_seqlen + 255) // 256 * 256) * _fp32_bytes + rows = max_logits_bytes // (block_q * row_bytes) * block_q + if rows == 0: + raise ValueError( + f'DSA score budget {max_logits_bytes} bytes is smaller than the minimum ' + f'aligned allocation {block_q * row_bytes} bytes. Increase ' + 'LMDEPLOY_DSA_INDEXER_MAX_LOGITS_MB.') + return rows return max(1, max_logits_bytes // (max_kv_seqlen * _fp32_bytes)) diff --git a/lmdeploy/pytorch/config.py b/lmdeploy/pytorch/config.py index 5a2e1c80ae..bb36d8e54c 100644 --- a/lmdeploy/pytorch/config.py +++ b/lmdeploy/pytorch/config.py @@ -443,6 +443,8 @@ class ModelConfig: use_flash_mla: bool = False mla_kv_cache_dtype: str | None = None mla_index_topk: int | None = None + # Custom sparse indexers also need score memory without selecting NSA. + reserve_dsa_score_workspace: bool = False # dllm model_paradigm: str = 'ar' diff --git a/lmdeploy/pytorch/configurations/glm5_next.py b/lmdeploy/pytorch/configurations/glm5_next.py index fcfd6a2774..1cb8b2fbaa 100644 --- a/lmdeploy/pytorch/configurations/glm5_next.py +++ b/lmdeploy/pytorch/configurations/glm5_next.py @@ -239,6 +239,7 @@ def build(cls, hf_config, model_path: str | None = None, **kwargs): # DeepSeek-V3.2 token indexer selected by mla_index_topk. Keeping this # unset also preserves the BF16 latent MLA cache policy. config.mla_index_topk = None + config.reserve_dsa_score_workspace = True config.k_head_dim = text_config.kv_lora_rank # Reuse Qwen3.5's token ring for convolution; recurrent/KPool states # keep a complete checkpoint after each verified token. diff --git a/lmdeploy/pytorch/engine/executor/base.py b/lmdeploy/pytorch/engine/executor/base.py index f9ec06bbb7..4fe9d4851e 100644 --- a/lmdeploy/pytorch/engine/executor/base.py +++ b/lmdeploy/pytorch/engine/executor/base.py @@ -218,7 +218,8 @@ def _get_rank_cache_block_sizes(cache_block_sizes: list[_WorkerCachePlanSizes]) def _get_dsa_score_workspace_size(self) -> int: """Return the bounded sparse-indexer score workspace in bytes.""" - if getattr(self.model_config, 'mla_index_topk', None) is None: + if (getattr(self.model_config, 'mla_index_topk', None) is None + and not getattr(self.model_config, 'reserve_dsa_score_workspace', False)): return 0 return _envs.dsa_indexer_max_logits_mb * (1 << 20) From fc843a3061a459f80952f257a3377b976a7eac03 Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Mon, 28 Sep 2026 13:54:56 +0000 Subject: [PATCH 33/39] perf(pytorch): skip invalid KPool compression and broadcast prefill groups --- lmdeploy/pytorch/backends/cuda/kpool.py | 15 ++++++++++++--- lmdeploy/pytorch/kernels/cuda/kpool.py | 21 +++++++++++++++------ lmdeploy/pytorch/models/glm5_next.py | 15 ++++++++++++--- 3 files changed, 39 insertions(+), 12 deletions(-) diff --git a/lmdeploy/pytorch/backends/cuda/kpool.py b/lmdeploy/pytorch/backends/cuda/kpool.py index 872a99b408..c75aaf9639 100644 --- a/lmdeploy/pytorch/backends/cuda/kpool.py +++ b/lmdeploy/pytorch/backends/cuda/kpool.py @@ -54,7 +54,7 @@ def kpool_decode_update_cuda(keys, scores, tail_keys, tail_scores, state_ids, compress = (compress_kpool if torch.cuda.get_device_capability(keys.device)[0] >= 9 else kpool_compress_quantize_cuda) values, scales = compress( - closed_keys, closed_scores, ape, mode='decode', round_scale=round_scale) + closed_keys, closed_scores, ape, mode='decode', round_scale=round_scale, valid=valid) cache_keys, cache_scales = kpool_packed_cache_views(packed_cache, keys.size(-1)) fill_indexed_key_cache(values, scales, groups, valid, block_offsets, cache_keys, cache_scales, page_step=pool_size) @@ -106,9 +106,15 @@ def kpool_compress_quantize_cuda( *, mode: str, round_scale: bool, + valid: Tensor | None = None, ) -> tuple[Tensor, Tensor]: """Compress and quantize closed pools with LMDeploy's reusable semantics.""" + if valid is not None: + # The reference fallback still computes every row, but closed buffers + # are only initialized for valid pools by the sequence updater. + slot_k = torch.where(valid[:, None, None], slot_k, 0) + slot_score = torch.where(valid[:, None, None], slot_score, 0) pooled = kpool_compress(slot_k, slot_score, ape, mode=mode) return kpool_quantize_fp8( pooled, @@ -119,13 +125,14 @@ def kpool_compress_quantize_cuda( def kpool_select_prefill_cuda(query_fp8, query_weight, packed_cache, q_seqlens, kv_seqlens, block_offsets, kv_flatten_size, - pool_size, topk): + pool_size, topk, *, return_groups=False): """Score and select all ragged prefill requests without host length reads.""" _validate_query(query_fp8, query_weight) rows = query_fp8.size(0) if not rows: - return torch.empty((0, topk + pool_size - 1), device=query_fp8.device, dtype=torch.int32) + return torch.empty((0, topk // pool_size if return_groups else topk + pool_size - 1), + device=query_fp8.device, dtype=torch.int32) counts, starts, seq, lengths, query_starts, query_ends = kpool_prefill_metadata( q_seqlens, kv_seqlens, rows, pool_size) keys, scales = kpool_packed_cache_views(packed_cache, query_fp8.size(-1)) @@ -161,6 +168,8 @@ def kpool_select_prefill_cuda(query_fp8, query_weight, packed_cache, selected[row_slice].copy_(chunk) # Release before allocating the next score chunk, including its padding. del logits, chunk + if return_groups: + return selected return kpool_expand_groups_cuda(selected, lengths, pool_size, topk, seq_lens=seq) diff --git a/lmdeploy/pytorch/kernels/cuda/kpool.py b/lmdeploy/pytorch/kernels/cuda/kpool.py index cabea0abeb..aac5962dce 100644 --- a/lmdeploy/pytorch/kernels/cuda/kpool.py +++ b/lmdeploy/pytorch/kernels/cuda/kpool.py @@ -37,9 +37,9 @@ def _update_kpool_kernel( tail_s = tl.where(slot == position, score, tail_s) row = step * BATCH + request offset = (row * POOL + slot) * WIDTH + d - tl.store(ClosedKeys + offset, tail_k, d < WIDTH) - tl.store(ClosedScores + offset, tail_s, d < WIDTH) close = position == POOL - 1 + tl.store(ClosedKeys + offset, tail_k, valid & close & (d < WIDTH)) + tl.store(ClosedScores + offset, tail_s, valid & close & (d < WIDTH)) tl.store(GroupIds + row, (history + step) // POOL) tl.store(Valid + row, valid & close) tail_k = tl.where(close, 0, tail_k) @@ -280,9 +280,12 @@ def rotate_kpool_query(query: torch.Tensor) -> torch.Tensor: @triton.jit -def _compress_kpool_kernel(K, S, A, O, Scale, WIDTH: tl.constexpr, POOL: tl.constexpr, +def _compress_kpool_kernel(K, S, A, O, Scale, Valid, HAS_VALID: tl.constexpr, WIDTH: tl.constexpr, POOL: tl.constexpr, ONLINE: tl.constexpr, ROUND: tl.constexpr, LEVELS: tl.constexpr): row = tl.program_id(0) + if HAS_VALID: + if not tl.load(Valid + row): + return d = tl.arange(0, WIDTH) maximum = tl.full((WIDTH,), -float('inf'), tl.float32) denominator = tl.full((WIDTH,), 0, tl.float32) @@ -316,24 +319,30 @@ def _compress_kpool_kernel(K, S, A, O, Scale, WIDTH: tl.constexpr, POOL: tl.cons def compress_kpool(keys: torch.Tensor, scores: torch.Tensor, ape: torch.Tensor, - *, mode: str, round_scale: bool) -> tuple[torch.Tensor, torch.Tensor]: + *, mode: str, round_scale: bool, + valid: torch.Tensor | None = None) -> tuple[torch.Tensor, torch.Tensor]: """Fuse weighted pooling, BF16 Hadamard rotation and one-block FP8 quantization. The online/two-pass reduction order and both BF16 round trips match the reference KPool operation. Precise libdevice - functions and disabled FMA fusion preserve its FP32 arithmetic boundaries. + functions and disabled FMA fusion preserve its FP32 arithmetic boundaries. Invalid rows are left unwritten and must + be masked by the cache writer. """ if mode not in ('extend', 'decode'): raise ValueError(f'Unsupported pool compression mode: {mode}') if keys.ndim != 3 or scores.shape != keys.shape or ape.shape != keys.shape[1:]: raise ValueError('Expected matching [groups, pool, width] keys/scores and [pool, width] APE.') groups, pool, width = keys.shape + if valid is not None and (valid.shape != (groups,) or valid.dtype != torch.bool or valid.device != keys.device): + raise ValueError('Expected a boolean validity mask with one entry per pool on the same device.') + if valid is not None: + valid = valid.contiguous() if width <= 0 or width & (width - 1): raise ValueError('Pool width must be a positive power of two.') out = torch.empty(groups, width, device=keys.device, dtype=torch.float8_e4m3fn) scale = torch.empty(groups, 1, device=keys.device, dtype=torch.float32) if groups: _compress_kpool_kernel[(groups,)]( - keys.contiguous(), scores.contiguous(), ape.contiguous(), out, scale, width, pool, + keys.contiguous(), scores.contiguous(), ape.contiguous(), out, scale, valid, valid is not None, width, pool, mode == 'extend', round_scale, width.bit_length() - 1, num_warps=4, enable_fp_fusion=False) return out, scale diff --git a/lmdeploy/pytorch/models/glm5_next.py b/lmdeploy/pytorch/models/glm5_next.py index d690499574..7973b8479e 100644 --- a/lmdeploy/pytorch/models/glm5_next.py +++ b/lmdeploy/pytorch/models/glm5_next.py @@ -19,6 +19,7 @@ kpool_decode_update_cuda, kpool_dense_indices_cuda, kpool_expand_groups_cuda, + kpool_prefill_metadata, kpool_prefill_update_cuda, kpool_rotate_query_cuda, kpool_score_contiguous_cuda, @@ -1056,6 +1057,7 @@ def _select_kpool_indices( is_owner = tp_group.rank == 0 total_rows = hidden_states.size(1) output_width = self.index_topk + self.index_kpool - 1 + broadcast_groups = attn_metadata.is_decoding or (hidden_states.is_cuda and dist_ctx.dist_config.attn_tp > 1) if attn_metadata.is_decoding: batch_size = attn_metadata.kv_seqlens.numel() steps = total_rows // batch_size @@ -1096,11 +1098,12 @@ def _select_kpool_indices( query_weight, indexer_k_cache, attn_metadata, + return_groups=broadcast_groups, ) else: logical_indices = torch.empty( total_rows, - self.index_topk // self.index_kpool if attn_metadata.is_decoding else output_width, + self.index_topk // self.index_kpool if broadcast_groups else output_width, dtype=torch.int32, device=hidden_states.device, ) @@ -1109,7 +1112,11 @@ def _select_kpool_indices( group = tp_group.gpu_group source_rank = dist_ctx.rank - tp_group.rank dist.broadcast(logical_indices, src=source_rank, group=group) - if attn_metadata.is_decoding: + if broadcast_groups: + if not attn_metadata.is_decoding: + _, _, seq_lens, group_lengths, _, _ = kpool_prefill_metadata( + attn_metadata.q_seqlens, attn_metadata.kv_seqlens, + total_rows, self.index_kpool) expand = kpool_expand_groups_cuda if logical_indices.is_cuda else kpool_expand_selected_groups logical_indices = expand(logical_indices, group_lengths, self.index_kpool, self.index_topk, seq_lens=seq_lens) @@ -1121,6 +1128,8 @@ def _select_kpool_indices_prefill( query_weight: torch.Tensor, indexer_k_cache: torch.Tensor, attn_metadata: Any, + *, + return_groups: bool = False, ) -> torch.Tensor: """Select request-local pooled history for chunked prefill.""" if query_fp8.is_cuda: @@ -1128,7 +1137,7 @@ def _select_kpool_indices_prefill( query_fp8, query_weight, indexer_k_cache, attn_metadata.q_seqlens, attn_metadata.kv_seqlens, attn_metadata.block_offsets, attn_metadata.kv_flatten_size, - self.index_kpool, self.index_topk) + self.index_kpool, self.index_topk, return_groups=return_groups) q_seqlens = attn_metadata.q_seqlens.tolist() kv_seqlens = attn_metadata.kv_seqlens.tolist() logical_parts = [] From ba4787d5c6eb4e6ecab83036f6eba56beff9e772 Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Mon, 28 Sep 2026 14:32:56 +0000 Subject: [PATCH 34/39] perf: fuse sparse MLA prefill index remapping and padding --- .../backends/cuda/attention/sparse_mla.py | 23 +++++++++++-------- 1 file changed, 14 insertions(+), 9 deletions(-) diff --git a/lmdeploy/pytorch/backends/cuda/attention/sparse_mla.py b/lmdeploy/pytorch/backends/cuda/attention/sparse_mla.py index 7f488d2770..cb13e04999 100644 --- a/lmdeploy/pytorch/backends/cuda/attention/sparse_mla.py +++ b/lmdeploy/pytorch/backends/cuda/attention/sparse_mla.py @@ -73,22 +73,27 @@ def map_strided_decode(self, indices: torch.Tensor, block_offsets: torch.Tensor, block_stride, token_stride, index_stride) def _map_flat_prefill_impl(self, indices: torch.Tensor, q_seqlens: torch.Tensor, - cu_seqlens_k: torch.Tensor): + cu_seqlens_k: torch.Tensor, index_alignment: int = 1): """Map request-local prefill indices into the flattened KV buffer.""" num_tokens = indices.size(0) - kv_offsets = torch.repeat_interleave(cu_seqlens_k[:-1], q_seqlens, output_size=num_tokens) - invalid = indices < 0 - indices = indices + kv_offsets[:, None] - indices[invalid] = -1 + if q_seqlens.numel() == 1: + kv_offsets = cu_seqlens_k[:1] + else: + kv_offsets = torch.repeat_interleave(cu_seqlens_k[:-1], q_seqlens, output_size=num_tokens) + indices = torch.where(indices < 0, -1, indices + kv_offsets[:, None]) + # Compile remapping and the sparse kernel's index-tile padding together. + padding = -indices.size(-1) % index_alignment + if padding: + indices = torch.nn.functional.pad(indices, (0, padding), value=-1) return indices[:, None] def map_flat_prefill(self, indices: torch.Tensor, q_seqlens: torch.Tensor, - cu_seqlens_k: torch.Tensor): + cu_seqlens_k: torch.Tensor, index_alignment: int = 1): """Map request-local prefill indices into the flattened KV buffer.""" if self._map_prefill_func is None: self._map_prefill_func = _try_dynamic_compile(self._map_flat_prefill_impl, - indices, q_seqlens, cu_seqlens_k) - return self._map_prefill_func(indices, q_seqlens, cu_seqlens_k) + indices, q_seqlens, cu_seqlens_k, index_alignment) + return self._map_prefill_func(indices, q_seqlens, cu_seqlens_k, index_alignment) @staticmethod @functools.cache @@ -155,7 +160,7 @@ def _prefill_sparse(self, query: torch.Tensor, flatten_k: torch.Tensor, """Run sparse prefill over flattened BF16 KV.""" indices = self.index_mapper.map_flat_prefill(nsa_indices, attn_metadata.q_seqlens, - attn_metadata.cu_seqlens_k) + attn_metadata.cu_seqlens_k, index_alignment=128) return self._flash_mla_sparse_forward(query, flatten_k, indices) def _decoding_sparse_bf16(self, query: torch.Tensor, k_cache: torch.Tensor, From 19307597e8036f0adb43a18e0fdb736f5e88c92b Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Tue, 29 Sep 2026 05:22:05 +0000 Subject: [PATCH 35/39] perf: vectorize mHC input conversion for long prefills Compile the FP32 input conversion for contiguous HC inputs with at least 8192 token rows. Preserve the existing FP32 GEMM and RMS statistics, rounding boundaries, and short-request execution path. H200 TP4/BS1/MTP5, output 5, 30 paired requests: 8K TTFT falls from 446.23 to 433.35 ms and total latency from 475.91 to 462.68 ms (2.78%). Full-model generated tokens and captured logits remain bitwise equal. 28 numeric/CUDA Graph checks pass; 2K shows no established speedup. --- lmdeploy/pytorch/nn/hc_prepost.py | 10 +++++++++- 1 file changed, 9 insertions(+), 1 deletion(-) diff --git a/lmdeploy/pytorch/nn/hc_prepost.py b/lmdeploy/pytorch/nn/hc_prepost.py index 5b31db8df7..52d1604141 100644 --- a/lmdeploy/pytorch/nn/hc_prepost.py +++ b/lmdeploy/pytorch/nn/hc_prepost.py @@ -8,6 +8,12 @@ from lmdeploy.pytorch.models.patch import get_build_model_context +@torch.compile(dynamic=True) +def _cast_fp32(x: torch.Tensor) -> torch.Tensor: + """Vectorize the large HC input conversion without changing reductions.""" + return x.float() + + class HcPrePost(nn.Module): """DeepSeek-V4 hyper-connection pre/post reduction wrapper.""" @@ -30,7 +36,9 @@ def pre( ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: from lmdeploy.pytorch.nn.norm import rms_scale hidden_states, dtype = x, x.dtype - x = x.flatten(2).float() + x = x.flatten(2) + # Long prefills amortize the additional compiled-call overhead. + x = _cast_fp32(x) if x.is_contiguous() and x.size(0) * x.size(1) >= 8192 else x.float() if self.avoid_gemv and x.size(0) == 1 and x.size(1) == 1: # Single-token decode otherwise selects GEMV, whose reduction # order can differ from multi-token speculative verification. From 932285b05f5792c71c25df56a5961b4c1c7438d2 Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Tue, 29 Sep 2026 05:28:02 +0000 Subject: [PATCH 36/39] style: wrap MoE reduction docstring for lint --- lmdeploy/pytorch/kernels/cuda/moe/fused_moe.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/lmdeploy/pytorch/kernels/cuda/moe/fused_moe.py b/lmdeploy/pytorch/kernels/cuda/moe/fused_moe.py index 75baa38734..c648325176 100644 --- a/lmdeploy/pytorch/kernels/cuda/moe/fused_moe.py +++ b/lmdeploy/pytorch/kernels/cuda/moe/fused_moe.py @@ -982,7 +982,8 @@ def _moe_reduce_kernel( def moe_reduce(hidden_states: torch.Tensor, topk_weights: torch.Tensor, fp32_acc: bool = False, *, output_scale: float = 1.0) -> torch.Tensor: - """Weight and reduce experts with optional FP32 products and output scaling.""" + """Weight and reduce experts with optional FP32 products and output + scaling.""" assert hidden_states.dim() == 3 assert topk_weights.dim() == 2 assert hidden_states.size(0) == topk_weights.size(0) From ec2ec44806c32da8bd5c045bef1fa35fbd7faa00 Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Tue, 29 Sep 2026 14:04:41 +0000 Subject: [PATCH 37/39] perf(pytorch): reduce GLM decode projection and metadata overhead Batch KDA gate projections through the linear backend while preserving checkpoint loaders, TP shards and the original LoRA path. Reuse MTP key/score projections, fuse raw token-cache indexing, and share target KPool metadata within each eager or captured forward. Prepare the next mHC FP32 input in post-expand after the original BF16 rounding boundary. Keep FP32 GEMM, normalization and TP reductions intact. Validation: full-model 2K/8K token and logits parity; ragged/rejection cache replay; TP1/4/8 and DP4/EP4 loading; real LoRA TP4 checks; existing HC tests; paired serving A/B and separate four-rank Torch Profiler. --- lmdeploy/pytorch/backends/cuda/hc_prepost.py | 6 + lmdeploy/pytorch/backends/cuda/kpool.py | 49 ++++++- lmdeploy/pytorch/backends/default/linear.py | 7 + lmdeploy/pytorch/backends/hc_prepost.py | 6 + lmdeploy/pytorch/backends/linear.py | 7 + .../pytorch/kernels/cuda/dsv4/hc_prepost.py | 11 ++ lmdeploy/pytorch/kernels/cuda/kpool.py | 120 ++++++++++++++++ lmdeploy/pytorch/models/glm5_next.py | 134 +++++++++++++----- lmdeploy/pytorch/nn/hc_prepost.py | 12 +- lmdeploy/pytorch/nn/linear/__init__.py | 9 +- lmdeploy/pytorch/nn/linear/default.py | 35 +++++ 11 files changed, 349 insertions(+), 47 deletions(-) diff --git a/lmdeploy/pytorch/backends/cuda/hc_prepost.py b/lmdeploy/pytorch/backends/cuda/hc_prepost.py index c9b81070c7..fbe0809809 100644 --- a/lmdeploy/pytorch/backends/cuda/hc_prepost.py +++ b/lmdeploy/pytorch/backends/cuda/hc_prepost.py @@ -37,3 +37,9 @@ def pre_reduce(self, x: torch.Tensor, pre: torch.Tensor, out_dtype: torch.dtype) def post_expand(self, x: torch.Tensor, residual: torch.Tensor, post: torch.Tensor, comb: torch.Tensor) -> torch.Tensor: return hc_post_expand(x, residual, post, comb, self.hc_mult) + + def post_expand_with_fp32(self, x: torch.Tensor, residual: torch.Tensor, post: torch.Tensor, + comb: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + fp32 = torch.empty((*x.shape[:-1], self.hc_mult, x.size(-1)), device=x.device, dtype=torch.float32) + out = hc_post_expand(x, residual, post, comb, self.hc_mult, out_fp32=fp32) + return out, fp32 diff --git a/lmdeploy/pytorch/backends/cuda/kpool.py b/lmdeploy/pytorch/backends/cuda/kpool.py index c75aaf9639..4b31ce685f 100644 --- a/lmdeploy/pytorch/backends/cuda/kpool.py +++ b/lmdeploy/pytorch/backends/cuda/kpool.py @@ -4,6 +4,7 @@ from __future__ import annotations import functools +from dataclasses import dataclass import torch from torch import Tensor @@ -13,10 +14,13 @@ from lmdeploy.pytorch.kernels.cuda.flatten_kv_cache import flatten_kv_cache from lmdeploy.pytorch.kernels.cuda.kpool import ( compress_kpool, + gather_kpool_token_tail, kpool_prefill_metadata, partition_kpool, + prepare_kpool_decode_metadata, rotate_kpool_query, update_kpool, + write_kpool_token_cache, ) from lmdeploy.pytorch.kernels.cuda.sparse_index_topk import ( is_sparse_index_topk_supported, @@ -36,6 +40,32 @@ # [tokens, topk] masks and int64 temporaries separately during prefill or decode. kpool_expand_groups_cuda = torch.compile(kpool_expand_selected_groups, dynamic=True, fullgraph=True) kpool_rotate_query_cuda = rotate_kpool_query +kpool_gather_token_tail_cuda = gather_kpool_token_tail +kpool_write_token_cache_cuda = write_kpool_token_cache + + +@dataclass(frozen=True) +class KPoolDecodeMetadata: + """Layer-independent metadata owned by one eager or captured forward.""" + + seq_lens: Tensor + group_lengths: Tensor + context_lens: Tensor + block_table: Tensor + schedule: Tensor | None + + +def kpool_decode_metadata_cuda(attn_metadata, rows, pool_size, with_scores=True): + """Prepare decode metadata once; graph replay refills the same buffers.""" + seq, groups, context, schedule_lengths, table = prepare_kpool_decode_metadata( + attn_metadata.q_seqlens, attn_metadata.kv_seqlens, + attn_metadata.block_offsets, rows, pool_size, with_scores) + schedule = None + if with_scores and table.size(1): + deep_gemm = _get_deep_gemm() + schedule = deep_gemm.get_paged_mqa_logits_metadata( + schedule_lengths, 64, deep_gemm.get_num_sms()) + return KPoolDecodeMetadata(seq, groups, context, table, schedule) def kpool_dense_indices_cuda(q_seqlens, kv_seqlens, rows, pool_size, topk): @@ -296,6 +326,8 @@ def kpool_score_paged_cuda( group_lengths: Tensor, pooled_block_offsets: Tensor, page_size: int = 64, + *, + metadata: KPoolDecodeMetadata | None = None, ) -> Tensor: """Score pooled decode history with DeepGEMM's paged MQA primitive.""" _validate_query(query_fp8, query_weight) @@ -320,12 +352,17 @@ def kpool_score_paged_cuda( (rows, 0), dtype=torch.float32, device=query_fp8.device) deep_gemm = _get_deep_gemm() - context_lens = group_lengths.to( - device=query_fp8.device, dtype=torch.int32).contiguous().view(-1, 1) - block_table = pooled_block_offsets.to( - device=query_fp8.device, dtype=torch.int32).contiguous() - schedule = deep_gemm.get_paged_mqa_logits_metadata( - context_lens.clamp(min=1), page_size, deep_gemm.get_num_sms()) + if metadata is None: + context_lens = group_lengths.to( + device=query_fp8.device, dtype=torch.int32).contiguous().view(-1, 1) + block_table = pooled_block_offsets.to( + device=query_fp8.device, dtype=torch.int32).contiguous() + schedule = deep_gemm.get_paged_mqa_logits_metadata( + context_lens.clamp(min=1), page_size, deep_gemm.get_num_sms()) + else: + context_lens = metadata.context_lens + block_table = metadata.block_table + schedule = metadata.schedule return deep_gemm.fp8_paged_mqa_logits( query_fp8.contiguous().unsqueeze(1), packed_cache, diff --git a/lmdeploy/pytorch/backends/default/linear.py b/lmdeploy/pytorch/backends/default/linear.py index e7eba90485..e72fcfe4d6 100644 --- a/lmdeploy/pytorch/backends/default/linear.py +++ b/lmdeploy/pytorch/backends/default/linear.py @@ -10,6 +10,13 @@ class DefaultLinearImpl(LinearImpl): """Linear implementation api.""" + supports_batched = True + + def forward_batched(self, x: torch.Tensor, weight: torch.Tensor) -> torch.Tensor: + """Batch equal-sized local projections without changing + accumulation.""" + return torch.bmm(x, weight.transpose(1, 2)) + def forward(self, x, weight: torch.Tensor, diff --git a/lmdeploy/pytorch/backends/hc_prepost.py b/lmdeploy/pytorch/backends/hc_prepost.py index ea0cd1c6ae..68e2cf6af6 100644 --- a/lmdeploy/pytorch/backends/hc_prepost.py +++ b/lmdeploy/pytorch/backends/hc_prepost.py @@ -37,6 +37,12 @@ def post_expand(self, x: torch.Tensor, residual: torch.Tensor, post: torch.Tenso """Expand one hidden state back to ``[..., hc, dim]``.""" raise NotImplementedError + def post_expand_with_fp32(self, x: torch.Tensor, residual: torch.Tensor, post: torch.Tensor, + comb: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + """Also prepare the rounded output for the next FP32 HC projection.""" + out = self.post_expand(x, residual, post, comb) + return out, out.float() + @dataclass(frozen=True) class HCPrePostBuildSpec(BuildSpec[HCPrePostImpl]): diff --git a/lmdeploy/pytorch/backends/linear.py b/lmdeploy/pytorch/backends/linear.py index 3e9cb71653..d630d484d9 100644 --- a/lmdeploy/pytorch/backends/linear.py +++ b/lmdeploy/pytorch/backends/linear.py @@ -11,6 +11,13 @@ class LinearImpl(ABC): """Linear implementation api.""" + supports_batched = False + + def forward_batched(self, x: torch.Tensor, weight: torch.Tensor) -> torch.Tensor: + """Apply independent local projections with weights [batch, out, + in].""" + raise NotImplementedError + def update_weights(self, weight: torch.Tensor, bias: torch.Tensor | None = None): """Update weights.""" return weight, bias diff --git a/lmdeploy/pytorch/kernels/cuda/dsv4/hc_prepost.py b/lmdeploy/pytorch/kernels/cuda/dsv4/hc_prepost.py index 9f88dba564..693b5d915c 100644 --- a/lmdeploy/pytorch/kernels/cuda/dsv4/hc_prepost.py +++ b/lmdeploy/pytorch/kernels/cuda/dsv4/hc_prepost.py @@ -101,6 +101,7 @@ def _hc_post_expand_kernel( post_ptr, comb_ptr, out_ptr, + fp32_ptr, x_stride_n, x_stride_d, residual_stride_n, @@ -117,6 +118,7 @@ def _hc_post_expand_kernel( dim: tl.constexpr, hc_mult: tl.constexpr, BLOCK_D: tl.constexpr, + STORE_FP32: tl.constexpr, ): row_h = tl.program_id(0) row_id = row_h // hc_mult @@ -141,6 +143,10 @@ def _hc_post_expand_kernel( acc += weight * residual tl.store(out_ptr + row_id * out_stride_n + out_h * out_stride_h + offs_d * out_stride_d, acc, mask=mask) + if STORE_FP32: + # Retain the original output rounding before preparing the next GEMM. + rounded = acc.to(out_ptr.dtype.element_ty).to(tl.float32) + tl.store(fp32_ptr + (row_id * hc_mult + out_h) * dim + offs_d, rounded, mask=mask) def hc_pre_reduce( @@ -186,11 +192,14 @@ def hc_post_expand( post: torch.Tensor, comb: torch.Tensor, hc_mult: int, + out_fp32: torch.Tensor | None = None, ) -> torch.Tensor: """Expand DeepSeek-V4 HC states from ``[..., dim]`` to ``[..., hc, dim]``.""" dim = x.size(-1) out_shape = (*x.shape[:-1], hc_mult, dim) + if out_fp32 is not None: + assert out_fp32.shape == out_shape and out_fp32.dtype == torch.float32 and out_fp32.is_contiguous() out = torch.empty(out_shape, device=x.device, dtype=x.dtype) if x.numel() == 0: return out @@ -213,6 +222,7 @@ def hc_post_expand( post, comb, out, + out_fp32, *x.stride(), *residual.stride(), *post.stride(), @@ -224,6 +234,7 @@ def hc_post_expand( dim, hc_mult, block_d, + out_fp32 is not None, num_warps=4, ) return out.reshape(out_shape) diff --git a/lmdeploy/pytorch/kernels/cuda/kpool.py b/lmdeploy/pytorch/kernels/cuda/kpool.py index aac5962dce..8e6a0f050d 100644 --- a/lmdeploy/pytorch/kernels/cuda/kpool.py +++ b/lmdeploy/pytorch/kernels/cuda/kpool.py @@ -346,3 +346,123 @@ def compress_kpool(keys: torch.Tensor, scores: torch.Tensor, ape: torch.Tensor, keys.contiguous(), scores.contiguous(), ape.contiguous(), out, scale, valid, valid is not None, width, pool, mode == 'extend', round_scale, width.bit_length() - 1, num_warps=4, enable_fp_fusion=False) return out, scale + + +@triton.jit +def _gather_token_tail_kernel(Cache, Blocks, Q, KV, Keys, Scores, StateIds, + stride_cb: tl.constexpr, stride_ct: tl.constexpr, + stride_cs: tl.constexpr, stride_cd: tl.constexpr, + stride_bb: tl.constexpr, stride_bp: tl.constexpr, + PAGES: tl.constexpr, PAGE_SIZE: tl.constexpr, + POOL: tl.constexpr, WIDTH: tl.constexpr, + BLOCK_D: tl.constexpr): + request = tl.program_id(0) + history = tl.load(KV + request).to(tl.int64) - tl.load(Q + request).to(tl.int64) + tail = history % POOL + slot = tl.arange(0, POOL) + position = history - tail + slot + page = position // PAGE_SIZE + valid = (slot < tail) & (position >= 0) & (page < PAGES) + block = tl.load(Blocks + request * stride_bb + page * stride_bp, valid, other=0) + dim = tl.arange(0, BLOCK_D) + ptr = Cache + block[:, None] * stride_cb + (position % PAGE_SIZE)[:, None] * stride_ct + ptr += dim[None, :] * stride_cd + mask = valid[:, None] & (dim[None, :] < WIDTH) + keys = tl.load(ptr, mask, other=0) + scores = tl.load(ptr + stride_cs, mask, other=0) + out = (request * POOL + slot[:, None]) * WIDTH + dim[None, :] + tl.store(Keys + out, keys, dim[None, :] < WIDTH) + tl.store(Scores + out, scores, dim[None, :] < WIDTH) + tl.store(StateIds + request, request) + + +def gather_kpool_token_tail(cache, block_offsets, q_seqlens, kv_seqlens, pool_size): + """Reconstruct private draft tails from pageable accepted token history.""" + batch, width = q_seqlens.numel(), cache.size(-1) + keys = cache.new_empty((batch, pool_size, width)) + scores = torch.empty_like(keys) + state_ids = torch.empty(batch, device=cache.device, dtype=torch.int64) + _gather_token_tail_kernel[(batch,)]( + cache, block_offsets, q_seqlens, kv_seqlens, keys, scores, state_ids, + *cache.stride(), *block_offsets.stride(), block_offsets.size(1), cache.size(1), + pool_size, width, triton.next_power_of_2(width), num_warps=4) + return keys, scores, state_ids + + +@triton.jit +def _write_token_cache_kernel(Cache, Blocks, Q, KV, Starts, Keys, Scores, + stride_cb: tl.constexpr, stride_ct: tl.constexpr, + stride_cs: tl.constexpr, stride_cd: tl.constexpr, + stride_bb: tl.constexpr, stride_bp: tl.constexpr, + stride_kt: tl.constexpr, stride_kd: tl.constexpr, + stride_st: tl.constexpr, stride_sd: tl.constexpr, + BATCH: tl.constexpr, PAGES: tl.constexpr, + PAGE_SIZE: tl.constexpr, WIDTH: tl.constexpr, + BLOCK_B: tl.constexpr, BLOCK_D: tl.constexpr): + row = tl.program_id(0) + request_ids = tl.arange(0, BLOCK_B) + starts = tl.load(Starts + request_ids, request_ids < BATCH, other=2147483647) + request = tl.sum((row >= starts).to(tl.int32), 0) - 1 + history = tl.load(KV + request).to(tl.int64) - tl.load(Q + request).to(tl.int64) + position = history + row - tl.load(Starts + request) + page = position // PAGE_SIZE + valid = (position >= 0) & (page < PAGES) + block = tl.load(Blocks + request * stride_bb + page * stride_bp, valid, other=0) + dim = tl.arange(0, BLOCK_D) + keys = tl.load(Keys + row * stride_kt + dim * stride_kd, dim < WIDTH, other=0) + scores = tl.load(Scores + row * stride_st + dim * stride_sd, dim < WIDTH, other=0) + ptr = Cache + block * stride_cb + (position % PAGE_SIZE) * stride_ct + dim * stride_cd + tl.store(ptr, keys, valid & (dim < WIDTH)) + tl.store(ptr + stride_cs, scores, valid & (dim < WIDTH)) + + +def write_kpool_token_cache(cache, keys, scores, block_offsets, q_seqlens, + kv_seqlens, cu_seqlens_q): + """Write raw draft projections without materializing token index + tensors.""" + batch, width = q_seqlens.numel(), cache.size(-1) + _write_token_cache_kernel[(keys.size(0),)]( + cache, block_offsets, q_seqlens, kv_seqlens, cu_seqlens_q, keys, scores, + *cache.stride(), *block_offsets.stride(), *keys.stride(), *scores.stride(), + batch, block_offsets.size(1), cache.size(1), width, + triton.next_power_of_2(batch), triton.next_power_of_2(width), num_warps=4) + + +@triton.jit +def _decode_metadata_kernel(Q, KV, Blocks, Seq, Groups, Context, ScheduleLengths, Table, + STEPS: tl.constexpr, POOL: tl.constexpr, PAGES: tl.constexpr, + stride_bb: tl.constexpr, stride_bp: tl.constexpr, + BLOCK: tl.constexpr): + row = tl.program_id(0) + tile = tl.program_id(1) + request, step = row // STEPS, row % STEPS + length = tl.load(KV + request).to(tl.int64) - tl.load(Q + request).to(tl.int64) + step + 1 + # Triton integer division truncates; match Torch floor division for padded rows. + groups = (length - tl.where(length < 0, POOL - 1, 0)) // POOL + if tile == 0: + tl.store(Seq + row, length) + tl.store(Groups + row, groups) + tl.store(Context + row, groups) + tl.store(ScheduleLengths + row, tl.maximum(groups, 1)) + pages = tile * BLOCK + tl.arange(0, BLOCK) + block = tl.load(Blocks + request * stride_bb + pages * POOL * stride_bp, + pages < PAGES, other=0) + tl.store(Table + row * PAGES + pages, block, pages < PAGES) + + +def prepare_kpool_decode_metadata(q_seqlens, kv_seqlens, block_offsets, rows, + pool_size, with_scores=True): + """Fill graph-owned sequence and pooled-page metadata once per forward.""" + batch = q_seqlens.numel() + if not batch or rows % batch: + raise ValueError('Decode rows must be a multiple of the request count.') + pages = triton.cdiv(block_offsets.size(1), pool_size) if with_scores else 0 + seq = torch.empty(rows, device=q_seqlens.device, dtype=torch.int64) + groups = torch.empty_like(seq) + context = torch.empty((rows, 1), device=q_seqlens.device, dtype=torch.int32) + schedule_lengths = torch.empty_like(context) + table = torch.empty((rows, pages), device=q_seqlens.device, dtype=torch.int32) + _decode_metadata_kernel[(rows, max(1, triton.cdiv(pages, 256)))]( + q_seqlens, kv_seqlens, block_offsets, seq, groups, context, schedule_lengths, + table, rows // batch, pool_size, pages, *block_offsets.stride(), 256, num_warps=4) + return seq, groups, context, schedule_lengths, table diff --git a/lmdeploy/pytorch/models/glm5_next.py b/lmdeploy/pytorch/models/glm5_next.py index 7973b8479e..176e1946a4 100644 --- a/lmdeploy/pytorch/models/glm5_next.py +++ b/lmdeploy/pytorch/models/glm5_next.py @@ -16,9 +16,11 @@ from lmdeploy.pytorch.backends.cuda.attention.sparse_mla import FlashMLASparseImpl from lmdeploy.pytorch.backends.cuda.kpool import ( kpool_compress_quantize_cuda, + kpool_decode_metadata_cuda, kpool_decode_update_cuda, kpool_dense_indices_cuda, kpool_expand_groups_cuda, + kpool_gather_token_tail_cuda, kpool_prefill_metadata, kpool_prefill_update_cuda, kpool_rotate_query_cuda, @@ -26,6 +28,7 @@ kpool_score_paged_cuda, kpool_select_groups_cuda, kpool_select_prefill_cuda, + kpool_write_token_cache_cuda, ) from lmdeploy.pytorch.configurations.glm5_next import is_glm5_kda_layer from lmdeploy.pytorch.consts import ( @@ -59,6 +62,7 @@ kpool_write_packed_cache_batched, ) from lmdeploy.pytorch.nn.linear import ( + build_batched_linear, build_colwise_linear, build_merged_colwise_linear, build_o_proj, @@ -646,9 +650,10 @@ def __init__(self, device=device, is_tp=True, ) - self.f_a_proj = build_colwise_linear( + # Both low-rank gate inputs are replicated across attention TP. + self.fg_a_proj = build_merged_colwise_linear( self.hidden_size, - self.head_dim, + [self.head_dim, self.head_dim], bias=False, quant_config=None, dtype=dtype, @@ -673,16 +678,7 @@ def __init__(self, device=device, is_tp=True, ) - self.g_a_proj = build_colwise_linear( - self.hidden_size, - self.head_dim, - bias=False, - quant_config=None, - dtype=dtype, - device=device, - is_tp=False, - ) - + self.fg_b_proj = build_batched_linear(self.f_b_proj, self.g_b_proj) self.qkv_conv1d = Glm5NextQKVConv1d(local_projection_size, self.conv_kernel_size, device=device) @@ -723,8 +719,7 @@ def forward(self, hidden_states: torch.Tensor, kda_metadata: GatedDeltaMeta) -> torch.Tensor: mixed_qkv = self.qkv_proj(hidden_states) raw_beta = self.b_proj(hidden_states) - raw_gate = self.f_b_proj(self.f_a_proj(hidden_states)) - norm_gate = self.g_b_proj(self.g_a_proj(hidden_states)) + raw_gate, norm_gate = self.fg_b_proj(self.fg_a_proj(hidden_states)) core_output = self.kda( mixed_qkv=mixed_qkv, @@ -875,6 +870,8 @@ def _update_kpool_cache( tail_state: Sequence[torch.Tensor], state_ids: torch.Tensor, attn_metadata: Any, + *, + projected: tuple[torch.Tensor, torch.Tensor] | None = None, ) -> torch.Tensor: """Compress closed pools and persist each request's unfinished tail.""" if tail_state is None or len(tail_state) != 2: @@ -886,8 +883,11 @@ def _update_kpool_cache( tail_k_state, tail_score_state = tail_state history_lengths = attn_metadata.kv_seqlens - attn_metadata.q_seqlens indexer_k_cache = self.indexer.get_block_cache() - key = self.indexer.project_key(hidden_states)[0] - score = self.indexer.project_compress_score(hidden_states)[0] + if projected is None: + key = self.indexer.project_key(hidden_states)[0] + score = self.indexer.project_compress_score(hidden_states)[0] + else: + key, score = projected if attn_metadata.is_decoding and key.is_cuda: batch_size = state_ids.numel() if not batch_size or key.size(0) % batch_size: @@ -1043,6 +1043,7 @@ def _select_kpool_indices( q_lora: torch.Tensor, indexer_k_cache: torch.Tensor, attn_metadata: Any, + kpool_metadata: Any = None, ) -> torch.Tensor: """Score/select on rank 0; expand decode group ids on each TP rank.""" if (hidden_states.is_cuda and not attn_metadata.is_decoding @@ -1059,12 +1060,19 @@ def _select_kpool_indices( output_width = self.index_topk + self.index_kpool - 1 broadcast_groups = attn_metadata.is_decoding or (hidden_states.is_cuda and dist_ctx.dist_config.attn_tp > 1) if attn_metadata.is_decoding: - batch_size = attn_metadata.kv_seqlens.numel() - steps = total_rows // batch_size - history = attn_metadata.kv_seqlens - attn_metadata.q_seqlens - step_ids = torch.arange(1, steps + 1, device=hidden_states.device) - seq_lens = (history[:, None] + step_ids).flatten().to(torch.int64) - group_lengths = torch.div(seq_lens, self.index_kpool, rounding_mode='floor') + if hidden_states.is_cuda: + if kpool_metadata is None: + kpool_metadata = kpool_decode_metadata_cuda( + attn_metadata, total_rows, self.index_kpool, is_owner) + seq_lens = kpool_metadata.seq_lens + group_lengths = kpool_metadata.group_lengths + else: + batch_size = attn_metadata.kv_seqlens.numel() + steps = total_rows // batch_size + history = attn_metadata.kv_seqlens - attn_metadata.q_seqlens + step_ids = torch.arange(1, steps + 1, device=hidden_states.device) + seq_lens = (history[:, None] + step_ids).flatten().to(torch.int64) + group_lengths = torch.div(seq_lens, self.index_kpool, rounding_mode='floor') if is_owner: query = self.indexer.project_query(q_lora)[0] @@ -1075,16 +1083,18 @@ def _select_kpool_indices( query_weight = (head_gate * query_scale.squeeze(-1) * self.indexer.softmax_scale) if attn_metadata.is_decoding: - pooled_block_offsets = kpool_pooled_block_offsets( - attn_metadata.block_offsets.repeat_interleave(steps, dim=0), - self.index_kpool, - ) + pooled_block_offsets = (kpool_metadata.block_table + if kpool_metadata is not None else + kpool_pooled_block_offsets( + attn_metadata.block_offsets.repeat_interleave(steps, dim=0), + self.index_kpool)) logits = kpool_score_paged_cuda( query_fp8, query_weight, indexer_k_cache, group_lengths, pooled_block_offsets, + metadata=kpool_metadata, ) selected_groups = kpool_select_groups_cuda( logits.contiguous(), @@ -1203,6 +1213,7 @@ def _kpool_indices( return_indices: bool, topk_indices_buffer: DSATopKIndicesBuffer | None = None, skip_topk: bool = False, + kpool_metadata: Any = None, ) -> torch.Tensor | None: indexer_k_cache = self._update_kpool_cache( hidden_states, tail_state, state_ids, attn_metadata) @@ -1212,7 +1223,7 @@ def _kpool_indices( if not return_indices and topk_indices_buffer is None: return None indices = self._select_kpool_indices( - hidden_states, q_lora, indexer_k_cache, attn_metadata) + hidden_states, q_lora, indexer_k_cache, attn_metadata, kpool_metadata) if topk_indices_buffer is not None: # MTP needs seed indices even for dense short prefill: subsequent # draft steps reuse its last-token rows through the shared proposer. @@ -1278,6 +1289,7 @@ def forward( state_ids: torch.Tensor | None = None, topk_indices_buffer: DSATopKIndicesBuffer | None = None, skip_topk: bool = False, + kpool_metadata: Any = None, ) -> torch.Tensor: dist_ctx = get_dist_manager().current_context() num_heads = self.num_heads // dist_ctx.dist_config.attn_tp @@ -1298,6 +1310,7 @@ def forward( return_indices=use_sparse, topk_indices_buffer=topk_indices_buffer, skip_topk=skip_topk, + kpool_metadata=kpool_metadata, ) if not use_sparse: return self._forward_prefill_mha( @@ -1337,6 +1350,7 @@ def forward( return_indices=True, topk_indices_buffer=topk_indices_buffer, skip_topk=skip_topk, + kpool_metadata=kpool_metadata, ) # NoPE queries and latent cache both contain exactly 512 values. query_states = self._absorbed_query(unabsorbed_query, num_heads) @@ -1437,15 +1451,25 @@ def __init__(self, requires_grad=False) def _hc_pre(self, hidden_states: torch.Tensor, fn: torch.Tensor, - scale: torch.Tensor, base: torch.Tensor, norm: RMSNorm): + scale: torch.Tensor, base: torch.Tensor, norm: RMSNorm, + x_fp32: torch.Tensor | None = None): return self.hc_prepost.pre( - hidden_states, fn, scale, base, norm.eps, norm_weight=norm.weight) + hidden_states, fn, scale, base, norm.eps, norm_weight=norm.weight, x_fp32=x_fp32) + + def _hc_post(self, hidden_states: torch.Tensor, residual: torch.Tensor, + post: torch.Tensor, comb: torch.Tensor, prepare_fp32: bool): + if prepare_fp32: + return self.hc_prepost.post_expand_with_fp32(hidden_states, residual, post, comb) + return self.hc_prepost.post_expand(hidden_states, residual, post, comb), None def forward(self, hidden_states: torch.Tensor, past_key_value: Sequence[torch.Tensor], attn_metadata: Any, kda_metadata: GatedDeltaMeta, kpool_tail_state: Sequence[torch.Tensor] | None = None, - state_ids: torch.Tensor | None = None) -> torch.Tensor: + state_ids: torch.Tensor | None = None, + kpool_metadata: Any = None, + hc_input_fp32: torch.Tensor | None = None, + prepare_hc_fp32: bool = False) -> tuple[torch.Tensor, torch.Tensor | None]: residual = hidden_states hidden_states, post, comb = self._hc_pre( hidden_states, @@ -1453,6 +1477,7 @@ def forward(self, hidden_states: torch.Tensor, self.hc_attn_scale, self.hc_attn_base, self.input_layernorm, + x_fp32=hc_input_fp32, ) if self.is_linear_attention: hidden_states = self.self_attn(hidden_states, @@ -1463,9 +1488,10 @@ def forward(self, hidden_states: torch.Tensor, past_key_value=past_key_value, attn_metadata=attn_metadata, kpool_tail_state=kpool_tail_state, - state_ids=state_ids) - hidden_states = self.hc_prepost.post_expand(hidden_states, residual, - post, comb) + state_ids=state_ids, + kpool_metadata=kpool_metadata) + hidden_states, hc_input_fp32 = self._hc_post(hidden_states, residual, + post, comb, prepare_hc_fp32) residual = hidden_states hidden_states, post, comb = self._hc_pre( @@ -1474,9 +1500,10 @@ def forward(self, hidden_states: torch.Tensor, self.hc_ffn_scale, self.hc_ffn_base, self.post_attention_layernorm, + x_fp32=hc_input_fp32, ) hidden_states = self.mlp(hidden_states) - return self.hc_prepost.post_expand(hidden_states, residual, post, comb) + return self._hc_post(hidden_states, residual, post, comb, prepare_hc_fp32) class Glm5NextModel(nn.Module): @@ -1537,18 +1564,28 @@ def forward( raise RuntimeError( f'GLM-5.3 expects {expected_full_layers} KPool tail rows, ' f'got {len(kpool_tail_states)}.') + kpool_metadata = None + if hidden_states.is_cuda and attn_metadata.is_decoding: + kpool_metadata = kpool_decode_metadata_cuda( + attn_metadata, hidden_states.size(1), self.config.index_kpool, + get_tp_world_rank('attn')[1] == 0) full_layer_row = 0 + hc_input_fp32 = None + prepare_hc_fp32 = hidden_states.is_cuda and hidden_states.size(0) * hidden_states.size(1) <= 16 for layer, past_key_value in zip(self.layers, past_key_values): kpool_tail_state = None if not layer.is_linear_attention: kpool_tail_state = kpool_tail_states[full_layer_row] full_layer_row += 1 - hidden_states = layer(hidden_states, + hidden_states, hc_input_fp32 = layer(hidden_states, past_key_value=past_key_value, attn_metadata=attn_metadata, kda_metadata=kda_metadata, kpool_tail_state=kpool_tail_state, - state_ids=state_ids) + state_ids=state_ids, + kpool_metadata=kpool_metadata, + hc_input_fp32=hc_input_fp32, + prepare_hc_fp32=prepare_hc_fp32) hidden_states = hidden_states.mean(dim=2) return self.norm(hidden_states) @@ -1559,6 +1596,10 @@ def get_input_embeddings(self): class Glm5NextForConditionalGeneration(DeepseekV32ForCausalLM): """GLM-5.3 conditional-generation wrapper for text, image and video.""" + packed_modules_mapping = { + 'fg_a_proj': ['f_a_proj', 'g_a_proj'], + } + def __init__(self, config: Any, ctx_mgr: StepContextManager, @@ -1866,6 +1907,8 @@ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]], *, ('.gate_up_proj', '.up_proj', 1), ] kda_params_mapping = [ + ('.fg_a_proj', '.f_a_proj', 0), + ('.fg_a_proj', '.g_a_proj', 1), ('.qkv_proj', '.q_proj', 'q'), ('.qkv_proj', '.k_proj', 'k'), ('.qkv_proj', '.v_proj', 'v'), @@ -1989,6 +2032,21 @@ def _update_kpool_cache(self, hidden_states, tail_state, state_ids, cache = (caches.row(binding.cache_name, binding.consumer_row) if hasattr(caches, 'row') else caches[binding.cache_name][binding.consumer_row]) + key = self.indexer.project_key(hidden_states)[0] + score = self.indexer.project_compress_score(hidden_states)[0] + if cache.is_cuda: + tail_keys, tail_scores, state_ids = kpool_gather_token_tail_cuda( + cache, attn_metadata.block_offsets, attn_metadata.q_seqlens, + attn_metadata.kv_seqlens, self.index_kpool) + result = super()._update_kpool_cache( + hidden_states, (tail_keys, tail_scores), state_ids, + attn_metadata, projected=(key, score)) + kpool_write_token_cache_cuda( + cache, key, score, attn_metadata.block_offsets, + attn_metadata.q_seqlens, attn_metadata.kv_seqlens, + attn_metadata.cu_seqlens_q) + return result + block_size = cache.size(1) history = (attn_metadata.kv_seqlens - attn_metadata.q_seqlens).long() tail_length = history.remainder(self.index_kpool) @@ -2001,7 +2059,7 @@ def _update_kpool_cache(self, hidden_states, tail_state, state_ids, state_ids = torch.arange(history.numel(), device=history.device) result = super()._update_kpool_cache( hidden_states, (tails[:, :, 0].contiguous(), tails[:, :, 1].contiguous()), - state_ids, attn_metadata) + state_ids, attn_metadata, projected=(key, score)) # Write raw projected tokens after reading the pre-forward tail. # Rejected positions are overwritten on their next visit. @@ -2011,8 +2069,6 @@ def _update_kpool_cache(self, hidden_states, tail_state, state_ids, token_ids = torch.arange(total_tokens, device=history.device) positions = history[batch] + token_ids - attn_metadata.cu_seqlens_q[batch] blocks = block_offsets[batch, positions.div(block_size, rounding_mode='floor')] - key = self.indexer.project_key(hidden_states)[0] - score = self.indexer.project_compress_score(hidden_states)[0] cache[blocks, positions.remainder(block_size)] = torch.stack((key, score), dim=1) return result diff --git a/lmdeploy/pytorch/nn/hc_prepost.py b/lmdeploy/pytorch/nn/hc_prepost.py index 52d1604141..f6feaf6ecb 100644 --- a/lmdeploy/pytorch/nn/hc_prepost.py +++ b/lmdeploy/pytorch/nn/hc_prepost.py @@ -33,12 +33,18 @@ def pre( hc_base: torch.Tensor, norm_eps: float, norm_weight: torch.Tensor | None = None, + x_fp32: torch.Tensor | None = None, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: from lmdeploy.pytorch.nn.norm import rms_scale hidden_states, dtype = x, x.dtype + if x_fp32 is not None: + assert x_fp32.shape == x.shape and x_fp32.dtype == torch.float32 x = x.flatten(2) # Long prefills amortize the additional compiled-call overhead. - x = _cast_fp32(x) if x.is_contiguous() and x.size(0) * x.size(1) >= 8192 else x.float() + if x_fp32 is not None: + x = x_fp32.flatten(2) + else: + x = _cast_fp32(x) if x.is_contiguous() and x.size(0) * x.size(1) >= 8192 else x.float() if self.avoid_gemv and x.size(0) == 1 and x.size(1) == 1: # Single-token decode otherwise selects GEMV, whose reduction # order can differ from multi-token speculative verification. @@ -55,3 +61,7 @@ def pre_reduce(self, x: torch.Tensor, pre: torch.Tensor, out_dtype: torch.dtype) def post_expand(self, x: torch.Tensor, residual: torch.Tensor, post: torch.Tensor, comb: torch.Tensor) -> torch.Tensor: return self.impl.post_expand(x, residual, post, comb) + + def post_expand_with_fp32(self, x: torch.Tensor, residual: torch.Tensor, post: torch.Tensor, + comb: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + return self.impl.post_expand_with_fp32(x, residual, post, comb) diff --git a/lmdeploy/pytorch/nn/linear/__init__.py b/lmdeploy/pytorch/nn/linear/__init__.py index 29ee45775e..cb502ae759 100644 --- a/lmdeploy/pytorch/nn/linear/__init__.py +++ b/lmdeploy/pytorch/nn/linear/__init__.py @@ -10,7 +10,7 @@ from .awq import AwqLinear, MergedAwqLinear, QKVAwqLinear from .blocked_fp8 import BlockedF8Linear, MergedBlockedF8Linear, QKVBlockedF8Linear -from .default import BaseLinear, MergedBaseLinear, QKVBaseLinear +from .default import BaseLinear, BatchedLinear, MergedBaseLinear, QKVBaseLinear from .lora import LoRA # noqa: F401 from .static_fp8 import ( MergedStaticF8Linear, @@ -20,6 +20,13 @@ from .w8a8 import MergedW8A8Linear, QKVW8A8Linear, StaticW8A8Linear, W8A8Linear +def build_batched_linear(*projections: nn.Module) -> nn.Module: + """Batch projections already registered by their parent module.""" + if not projections: + raise ValueError('At least one projection is required.') + return BatchedLinear(projections) + + def _is_static_per_tensor_fp8(quant_config): return ( quant_config.activation_scheme == 'static' diff --git a/lmdeploy/pytorch/nn/linear/default.py b/lmdeploy/pytorch/nn/linear/default.py index f959eb8025..60521e80e5 100644 --- a/lmdeploy/pytorch/nn/linear/default.py +++ b/lmdeploy/pytorch/nn/linear/default.py @@ -14,6 +14,41 @@ from .utils import QKVMixin, check_qkv_split_layout +class BatchedLinear(torch.nn.Module): + """Group projections while retaining their original owners and loaders.""" + + def __init__(self, projections): + super().__init__() + self.projections = tuple(projections) + first = self.projections[0] + self.in_features = first.in_features + self.out_features = first.out_features + if any((p.in_features, p.out_features) != (self.in_features, self.out_features) + for p in self.projections): + raise ValueError('Batched projections must have matching local dimensions.') + self.register_buffer('_weight', None, persistent=False) + + def process_weights_after_loading(self): + """Pack after backend weight updates; preserve individual + parameters.""" + self._weight = None + if all(type(p) is BaseLinear and p.colwise and not p.all_reduce + and p.tp_mode != TPMode.DP_TP and p.bias is None + and p.weight.is_cuda and p.weight.dtype == torch.bfloat16 + and p.impl.supports_batched for p in self.projections): + self._weight = torch.stack([p.get_unquantized_weight(torch.bfloat16) for p in self.projections]) + + def forward(self, x: torch.Tensor) -> tuple[torch.Tensor, ...]: + """Project concatenated inputs, retaining the original LoRA path.""" + if x.size(-1) != len(self.projections) * self.in_features: + raise ValueError('Batched input width must match the concatenated projection inputs.') + if self._weight is None or x.numel() == 0 or any(p.lora_adapters for p in self.projections): + return tuple(p(part) for p, part in zip(self.projections, x.split(self.in_features, dim=-1))) + inputs = x.reshape(-1, len(self.projections), self.in_features).transpose(0, 1) + output = self.projections[0].impl.forward_batched(inputs, self._weight) + return tuple(part.view(*x.shape[:-1], self.out_features) for part in output.unbind(0)) + + class BaseLinear(LinearBase): """Linear layer.""" From e7be6ba96a6d801ab61c022e6950a13f6265148f Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Wed, 30 Sep 2026 12:29:21 +0000 Subject: [PATCH 38/39] perf(pytorch): fuse masked GLM expert activation and FP8 quantization --- .../pytorch/backends/cuda/moe/blocked_fp8.py | 9 ++- lmdeploy/pytorch/kernels/cuda/activation.py | 81 +++++++++++++++---- lmdeploy/pytorch/models/glm5_next.py | 3 + 3 files changed, 77 insertions(+), 16 deletions(-) diff --git a/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py b/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py index 8f42e954a6..9cd934304e 100644 --- a/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py +++ b/lmdeploy/pytorch/backends/cuda/moe/blocked_fp8.py @@ -194,8 +194,13 @@ def experts( (gateup_output.shape[0], gateup_output.shape[1], gateup_output.shape[2] // 2 // self.block_size), device=gateup_output.device, dtype=torch.float32) - if act_func is None: - silu_and_mul_masked_post_quant_fwd(gateup_output, down_input, down_input_scale, self.block_size, masked_m) + # Custom activations can explicitly supply a masked FP8 fusion while + # retaining the existing callable fallback for other implementations. + masked_post_quant = (silu_and_mul_masked_post_quant_fwd if act_func is None + else getattr(act_func, 'masked_post_quant', None)) + if masked_post_quant is not None: + masked_post_quant(gateup_output, down_input, down_input_scale, self.block_size, masked_m, + scale_fmt=self.scale_fmt) else: # Only masked_m valid rows are consumed by the following GEMM. # Reuse the model's activation and the shared quantizer unchanged. diff --git a/lmdeploy/pytorch/kernels/cuda/activation.py b/lmdeploy/pytorch/kernels/cuda/activation.py index c198e4999e..5b92b4440c 100644 --- a/lmdeploy/pytorch/kernels/cuda/activation.py +++ b/lmdeploy/pytorch/kernels/cuda/activation.py @@ -4,6 +4,7 @@ import triton.language as tl from triton.language.extra import libdevice +from .blocked_gemm_fp8 import fast_round_scale from .utils import get_device_props fast_expf = tl.math.exp @@ -112,6 +113,7 @@ def _silu_and_mul_moe_ep_kernel( gateup_ptr, out_ptr, mask_ptr, + scale_ptr, N: tl.constexpr, M: tl.constexpr, stride_gue: tl.constexpr, @@ -122,6 +124,12 @@ def _silu_and_mul_moe_ep_kernel( stride_on: tl.constexpr, stride_m: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, + SWIGLU_LIMIT: tl.constexpr, + PRECISE_MUL: tl.constexpr, + GROUP_SIZE: tl.constexpr, + ROUND_SCALE: tl.constexpr, + FP8_MIN: tl.constexpr, + FP8_MAX: tl.constexpr, ): """Silu and mul kernel.""" n_block_id = tl.program_id(0) @@ -138,21 +146,48 @@ def _silu_and_mul_moe_ep_kernel( mask_m = tl.load(mask_ptr + e_id * stride_m) mask_m = tl.minimum(mask_m, M) - if mask_m < m_id_start: + if mask_m <= m_id_start: return gate_ptrs = gateup_ptr + e_id * stride_gue + m_id_start * stride_gum + offs_n * stride_gun up_ptrs = gate_ptrs + N * stride_gun out_ptrs = out_ptr + e_id * stride_oe + m_id_start * stride_om + offs_n * stride_on - for _ in tl.range(m_id_start, mask_m, m_id_stride): + for m_id in tl.range(m_id_start, mask_m, m_id_stride): gate = tl.load(gate_ptrs, mask=mask) up = tl.load(up_ptrs, mask=mask) + if SWIGLU_LIMIT is not None: + gate = tl.minimum(gate, SWIGLU_LIMIT) + up = tl.maximum(tl.minimum(up, SWIGLU_LIMIT), -SWIGLU_LIMIT) # exp expect fp32 gate = gate.to(tl.float32) - gate = gate / (1 + fast_expf(-gate)) - gate = gate.to(gateup_ptr.dtype.element_ty) + if PRECISE_MUL: + exp_neg_gate = libdevice.exp(-gate) + else: + exp_neg_gate = fast_expf(-gate) + gate = gate / (1 + exp_neg_gate) + if not PRECISE_MUL: + gate = gate.to(gateup_ptr.dtype.element_ty) out = gate * up + if scale_ptr is not None: + # Preserve the materialized activation's rounding before quantization. + out = out.to(gateup_ptr.dtype.element_ty) + groups: tl.constexpr = BLOCK_SIZE_N // GROUP_SIZE + values = out.reshape(groups, GROUP_SIZE) + amax = tl.max(tl.abs(values), axis=1) + amax = tl.maximum(amax, 1e-6).to(tl.float32) + if ROUND_SCALE: + scale = fast_round_scale(amax, 1 / FP8_MAX) + rscale = 1 / scale + else: + scale = amax * (1 / FP8_MAX) + rscale = FP8_MAX / amax + out = values.to(tl.float32) * rscale[:, None] + out = tl.clamp(out, FP8_MIN, FP8_MAX).reshape(BLOCK_SIZE_N) + group_offsets = n_block_id * groups + tl.arange(0, groups) + scale_offsets = (e_id * M + m_id) * (N // GROUP_SIZE) + group_offsets + tl.store(scale_ptr + scale_offsets, scale, group_offsets < N // GROUP_SIZE) + tl.store(out_ptrs, out, mask=mask) gate_ptrs += m_id_stride * stride_gum @@ -160,7 +195,10 @@ def _silu_and_mul_moe_ep_kernel( out_ptrs += m_id_stride * stride_om -def silu_and_mul_moe_ep(gate_up: torch.Tensor, mask_m: torch.Tensor, out: torch.Tensor = None): +def silu_and_mul_moe_ep(gate_up: torch.Tensor, mask_m: torch.Tensor, out: torch.Tensor = None, + swiglu_limit: float | None = None, precise_mul: bool = False, + output_scale: torch.Tensor = None, quant_group_size: int = 128, + scale_fmt: str | None = None): """Silu and mul for moe with expert parallelism.""" # gate_up: [num_experts, batch_size, 2*hidden_size] assert gate_up.dim() == 3 @@ -190,9 +228,17 @@ def silu_and_mul_moe_ep(gate_up: torch.Tensor, mask_m: torch.Tensor, out: torch. grid_size0 = triton.cdiv(N, BLOCK_SIZE_N) grid_size1 = min(M, triton.cdiv(ctas_per_device, grid_size0 * E)) grid = (grid_size0, E, grid_size1) + if output_scale is not None: + assert out.is_contiguous() and output_scale.is_contiguous() + assert out.shape == (E, M, N) + assert output_scale.shape == (E, M, N // quant_group_size) + assert N % quant_group_size == 0 and BLOCK_SIZE_N % quant_group_size == 0 + assert scale_fmt in (None, 'ue8m0') + finfo = torch.finfo(out.dtype) _silu_and_mul_moe_ep_kernel[grid](gate_up, out, mask_m, + output_scale, N, M, stride_gue=gate_up.stride(0), @@ -203,6 +249,12 @@ def silu_and_mul_moe_ep(gate_up: torch.Tensor, mask_m: torch.Tensor, out: torch. stride_on=out.stride(2), stride_m=mask_m.stride(0), BLOCK_SIZE_N=BLOCK_SIZE_N, + SWIGLU_LIMIT=swiglu_limit, + PRECISE_MUL=precise_mul, + GROUP_SIZE=quant_group_size, + ROUND_SCALE=scale_fmt == 'ue8m0', + FP8_MIN=finfo.min if output_scale is not None else 0, + FP8_MAX=finfo.max if output_scale is not None else 0, num_warps=num_warps, num_stages=num_stages) @@ -210,9 +262,13 @@ def silu_and_mul_moe_ep(gate_up: torch.Tensor, mask_m: torch.Tensor, out: torch. def silu_and_mul_masked_post_quant_fwd(input: torch.Tensor, output: torch.Tensor, output_scale: torch.Tensor, - quant_group_size: int, masked_m: torch.Tensor): - """Apply masked MoE SiLU-and-mul, then quantize to the preallocated FP8 - output.""" + quant_group_size: int, masked_m: torch.Tensor, + swiglu_limit: float | None = None, precise_mul: bool = False, + scale_fmt: str | None = None): + """Fuse activation and FP8 quantization for valid expert rows only. + + Invalid rows are left untouched; the following masked GEMM must ignore them. + """ assert input.is_contiguous() assert output.is_contiguous() assert input.dim() == 3 @@ -220,9 +276,6 @@ def silu_and_mul_masked_post_quant_fwd(input: torch.Tensor, output: torch.Tensor assert input.shape[-1] % 2 == 0 size_n = input.shape[-1] // 2 assert size_n % quant_group_size == 0 - activated = silu_and_mul_moe_ep(input, masked_m) - from .blocked_gemm_fp8 import _quant_fp8_launcher - _quant_fp8_launcher(activated.reshape(-1, size_n), - quant_group_size, - output.reshape(-1, size_n), - output_scale.reshape(-1, size_n // quant_group_size)) + silu_and_mul_moe_ep(input, masked_m, output, swiglu_limit=swiglu_limit, + precise_mul=precise_mul, output_scale=output_scale, + quant_group_size=quant_group_size, scale_fmt=scale_fmt) diff --git a/lmdeploy/pytorch/models/glm5_next.py b/lmdeploy/pytorch/models/glm5_next.py index 176e1946a4..ee4cef80ac 100644 --- a/lmdeploy/pytorch/models/glm5_next.py +++ b/lmdeploy/pytorch/models/glm5_next.py @@ -39,6 +39,7 @@ ) from lmdeploy.pytorch.distributed import get_dist_group, get_dist_manager, get_tp_world_rank from lmdeploy.pytorch.engine.cache_engine.schema import BlockCacheRequest +from lmdeploy.pytorch.kernels.cuda.activation import silu_and_mul_masked_post_quant_fwd from lmdeploy.pytorch.model_inputs import StepContext, StepContextManager, get_step_ctx_manager from lmdeploy.pytorch.nn import ( ApplyRotaryEmb, @@ -115,6 +116,8 @@ def _glm_swiglu_impl(intermediate: torch.Tensor, _GLM53_COMPACT_FP8_MOE_ACT = partial( _glm_swiglu_impl, swiglu_limit=10.0, precise_mul=True) +_GLM53_COMPACT_FP8_MOE_ACT.masked_post_quant = partial( + silu_and_mul_masked_post_quant_fwd, swiglu_limit=10.0, precise_mul=True) class Glm5NextVisionPatchEmbed(Glm4vVisionPatchEmbed): From ed1eea00129bf95ee35130e8aacaaa59e0868571 Mon Sep 17 00:00:00 2001 From: qescccczmr <105876751+qescccczmr@users.noreply.github.com> Date: Wed, 30 Sep 2026 12:56:54 +0000 Subject: [PATCH 39/39] perf: preserve padded KPool score strides for top-k --- lmdeploy/pytorch/models/glm5_next.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/lmdeploy/pytorch/models/glm5_next.py b/lmdeploy/pytorch/models/glm5_next.py index ee4cef80ac..3c72e4ca40 100644 --- a/lmdeploy/pytorch/models/glm5_next.py +++ b/lmdeploy/pytorch/models/glm5_next.py @@ -1100,7 +1100,7 @@ def _select_kpool_indices( metadata=kpool_metadata, ) selected_groups = kpool_select_groups_cuda( - logits.contiguous(), + logits, group_lengths, group_topk=self.index_topk // self.index_kpool, ) @@ -1190,7 +1190,7 @@ def _select_kpool_indices_prefill( ) group_budget = self.index_topk // self.index_kpool selected_groups = kpool_select_groups_cuda( - logits.contiguous(), + logits, group_lengths, group_topk=group_budget, max_group_length=num_groups,