diff --git a/lmdeploy/archs.py b/lmdeploy/archs.py index 2c51e91a50..b6e69e7e78 100644 --- a/lmdeploy/archs.py +++ b/lmdeploy/archs.py @@ -97,6 +97,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/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/attention.py b/lmdeploy/pytorch/backends/attention.py index 315c8f1943..7f6d458a29 100644 --- a/lmdeploy/pytorch/backends/attention.py +++ b/lmdeploy/pytorch/backends/attention.py @@ -190,6 +190,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/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/attention/mla.py b/lmdeploy/pytorch/backends/cuda/attention/mla.py index 7df5bc6a8e..57e55e4892 100644 --- a/lmdeploy/pytorch/backends/cuda/attention/mla.py +++ b/lmdeploy/pytorch/backends/cuda/attention/mla.py @@ -562,6 +562,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/sparse_mla.py b/lmdeploy/pytorch/backends/cuda/attention/sparse_mla.py index d0d7f56d9b..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 @@ -140,6 +145,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] @@ -149,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, 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/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/hc_prepost.py b/lmdeploy/pytorch/backends/cuda/hc_prepost.py index d3b8956b5a..fbe0809809 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: @@ -32,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/kda.py b/lmdeploy/pytorch/backends/cuda/kda.py new file mode 100644 index 0000000000..0ad3475a8b --- /dev/null +++ b/lmdeploy/pytorch/backends/cuda/kda.py @@ -0,0 +1,279 @@ +# 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 +state-ring kernel with gated-delta rule for both AR and MTP decode. +""" + +from copy import copy +from typing import Any + +import torch + +from lmdeploy.pytorch.backends.kda import KdaImpl + +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: + valid = metadata.valid_state + if metadata.is_init is not None: + 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: + ids = torch.where(metadata.valid_state, metadata.state_ids, -1) + _state_scatter(state.unsqueeze(1), ids, torch.zeros_like(ids), value) + + +class CudaKdaImpl(KdaImpl): + """FLA prefill and shared channelwise gated-delta decode.""" + + def __init__(self): + try: + from fla.modules.conv.triton.ops import ( + causal_conv1d_fwd, + causal_conv1d_update, + ) + 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.' + ) from exc + 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 + register_step_metadata_impl(self) + + 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): + """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 = torch.where(metadata.valid_state, metadata.state_ids, -1).long() + read_slot = history.remainder(ring_size) + # 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 + 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) + _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() + 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 + + def _decode_recurrent(self, q, k, v, g, beta, A_log, dt_bias, initial_state, + 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_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, + 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): + """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. + """ + 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.') + history = metadata.cache_seqlens + signed_ids = torch.where(metadata.valid_state, ids, -1) + values = mixed_qkv.reshape(batch, steps, -1).transpose(1, 2) + 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 = 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) + for x in mixed.split(heads * dim, dim=-1)] + 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( + 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 (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 + 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)) + 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() + 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)) + + 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.contiguous(), + k=k.contiguous(), + v=v.contiguous(), + 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 diff --git a/lmdeploy/pytorch/backends/cuda/kpool.py b/lmdeploy/pytorch/backends/cuda/kpool.py new file mode 100644 index 0000000000..4b31ce685f --- /dev/null +++ b/lmdeploy/pytorch/backends/cuda/kpool.py @@ -0,0 +1,375 @@ +# Copyright (c) OpenMMLab. All rights reserved. +"""LMDeploy CUDA adapters for pooled DSA selection.""" + +from __future__ import annotations + +import functools +from dataclasses import dataclass + +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 ( + 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, + 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 +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. +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): + """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, 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) + + +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 +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, + 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, + block_size=pooled.size(-1), + round_scale=round_scale, + ) + + +def kpool_select_prefill_cuda(query_fp8, query_weight, packed_cache, + q_seqlens, kv_seqlens, block_offsets, kv_flatten_size, + 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 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)) + 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 + 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 + if return_groups: + return selected + return kpool_expand_groups_cuda(selected, lengths, pool_size, topk, seq_lens=seq) + + +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, + q_seqlens, + lengths.clamp(max=max_group_length), + group_topk, + 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, + ) + + +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, + *, + metadata: KPoolDecodeMetadata | None = None, +) -> 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() + 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, + 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 94f0d4ec02..9cd934304e 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,20 @@ 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) + # 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. + 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 +226,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 +253,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 +279,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 +289,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 +303,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): @@ -289,13 +315,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, + output_scale: float = 1.0, + fp32_acc: bool = False): 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.output_scale = output_scale + self.fp32_acc = fp32_acc def ep_expert_list(self, world_size: int, rank: int): """Experts list of current rank.""" @@ -344,13 +374,17 @@ def forward(self, expert_offset=expert_offset, num_experts=num_experts, renormalize=self.renormalize, - act_func=act_func) + act_func=act_func, + output_scale=self.output_scale, + fp32_acc=self.fp32_acc) output = output.unflatten(0, input_size[:-1]) return output class FusedDeepEpMoEBlockedF8Impl(TritonFusedMoEBlockedF8Impl): + output_scale = 1.0 + def __init__(self, ep_size: int, ep_group: dist.ProcessGroup, @@ -362,8 +396,10 @@ 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, + 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 @@ -428,7 +464,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 @@ -518,7 +554,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, @@ -533,6 +572,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.fp32_acc, chunk_size=16 * 1024) return deepep_moe @@ -540,7 +580,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.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, @@ -553,6 +592,8 @@ 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, + fp32_acc=spec.fp32_acc, ) else: impl = TritonFusedMoEBlockedF8Impl( @@ -561,6 +602,8 @@ def _build_fused_moe_blocked_f8(spec: FusedMoEBlockedF8BuildSpec) -> FusedMoEBlo renormalize=spec.renormalize, 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 d0b1c31033..434f8a898e 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, + 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) @@ -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, + output_scale=self.output_scale, + fp32_acc=self.fp32_acc) # modify from dlblas: https://github.com/DeepLink-org/DLBlas @@ -375,11 +384,15 @@ def __init__( num_experts: int, hidden_dim: int, renormalize: 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, + fp32_acc: bool = False, ): - super().__init__(top_k, num_experts, renormalize) + 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 @@ -540,6 +553,8 @@ def _build_fused_moe(spec: FusedMoEBuildSpec) -> FusedMoEImpl: num_experts=spec.num_experts, 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, @@ -548,4 +563,6 @@ def _build_fused_moe(spec: FusedMoEBuildSpec) -> FusedMoEImpl: top_k=spec.top_k, num_experts=spec.num_experts, renormalize=spec.renormalize, + output_scale=spec.output_scale, + fp32_acc=spec.fp32_acc, ) 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/backends/cuda/op_backend.py b/lmdeploy/pytorch/backends/cuda/op_backend.py index 5847064776..3294b54f4c 100644 --- a/lmdeploy/pytorch/backends/cuda/op_backend.py +++ b/lmdeploy/pytorch/backends/cuda/op_backend.py @@ -43,6 +43,7 @@ def build_op(cls, spec: BuildSpec[ImplT], *, enable_deterministic: bool = False) from ..gated_delta_rule import GatedDeltaMetaBuildSpec, GatedDeltaRuleBuildSpec from ..hc_prepost import HCPrePostBuildSpec from ..indexer import V4IndexerBuildSpec + from ..kda import KdaBuildSpec from ..lora import LoRABuildSpec from ..moe import ( FusedMoEBlockedF8BuildSpec, @@ -62,9 +63,12 @@ def build_op(cls, spec: BuildSpec[ImplT], *, enable_deterministic: bool = False) if isinstance(spec, SiluAndMulBuildSpec): from .activation import TritonSiluAndMulImpl return cast(ImplT, TritonSiluAndMulImpl(spec.inplace)) + if isinstance(spec, KdaBuildSpec): + from .kda import CudaKdaImpl + 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/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/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/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/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/backends/hc_prepost.py b/lmdeploy/pytorch/backends/hc_prepost.py index fe68a598fd..68e2cf6af6 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 @@ -35,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/kda.py b/lmdeploy/pytorch/backends/kda.py new file mode 100644 index 0000000000..783d8b8b0a --- /dev/null +++ b/lmdeploy/pytorch/backends/kda.py @@ -0,0 +1,37 @@ +# Copyright (c) OpenMMLab. All rights reserved. +from abc import ABC, abstractmethod +from dataclasses import dataclass +from typing import Any + +import torch + +from .base import BuildSpec + + +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 + + +@dataclass(frozen=True) +class KdaBuildSpec(BuildSpec[KdaImpl]): + """Request a device-specific KDA implementation.""" 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/backends/moe.py b/lmdeploy/pytorch/backends/moe.py index 01bb13175c..810af52648 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 + output_scale: float = 1.0 + fp32_acc: bool = False class FusedMoEW8A8Impl(ABC): @@ -256,6 +258,8 @@ class FusedMoEBlockedF8BuildSpec(BuildSpec[FusedMoEBlockedF8Impl]): layer_idx: int 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/config.py b/lmdeploy/pytorch/config.py index 4a8a06cad5..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' @@ -460,6 +462,12 @@ class ModelConfig: state_cache_specs: list[StateCacheSpec] = field(default_factory=list) use_standard_kv_cache: bool = True + # 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 diff --git a/lmdeploy/pytorch/configurations/glm5_next.py b/lmdeploy/pytorch/configurations/glm5_next.py new file mode 100644 index 0000000000..1cb8b2fbaa --- /dev/null +++ b/lmdeploy/pytorch/configurations/glm5_next.py @@ -0,0 +1,298 @@ +# Copyright (c) OpenMMLab. All rights reserved. +"""PyTorch engine configuration for GLM-5.3-Flash.""" + +import torch + +from lmdeploy.pytorch.config import StateCacheSpec +from lmdeploy.pytorch.consts import ( + GLM5_KDA_CONV_STATE, + GLM5_KDA_RECURRENT_STATE, + GLM5_KPOOL_TAIL_K_STATE, + GLM5_KPOOL_TAIL_SCORE_STATE, +) + +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.') + + +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)) + 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, + **dict(kwargs, is_draft_model=False)) + + tp = kwargs.get('tp', 1) + 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.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. + 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_spec_tokens), + torch.bfloat16, + ), + StateCacheSpec( + GLM5_KDA_RECURRENT_STATE, + (num_linear_layers, *ring_shape, local_heads, head_dim, head_dim), + torch.float32, + ), + StateCacheSpec( + GLM5_KPOOL_TAIL_K_STATE, + (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, *ring_shape, 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 + 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: + hf_config.dtype = text_dtype + return config diff --git a/lmdeploy/pytorch/consts.py b/lmdeploy/pytorch/consts.py index a94ca379aa..0a08a272be 100644 --- a/lmdeploy/pytorch/consts.py +++ b/lmdeploy/pytorch/consts.py @@ -19,6 +19,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 026060683a..4fe9d4851e 100644 --- a/lmdeploy/pytorch/engine/executor/base.py +++ b/lmdeploy/pytorch/engine/executor/base.py @@ -62,6 +62,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.') @@ -212,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) diff --git a/lmdeploy/pytorch/engine/logits_process.py b/lmdeploy/pytorch/engine/logits_process.py index fb0024f992..30853ff882 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/activation.py b/lmdeploy/pytorch/kernels/cuda/activation.py index 6eebe91af7..5b92b4440c 100644 --- a/lmdeploy/pytorch/kernels/cuda/activation.py +++ b/lmdeploy/pytorch/kernels/cuda/activation.py @@ -2,7 +2,9 @@ import torch import triton 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 @@ -19,6 +21,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 +44,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 +66,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 +100,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) @@ -96,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, @@ -106,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) @@ -122,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 @@ -144,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 @@ -174,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), @@ -187,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) @@ -194,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 @@ -204,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/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/kernels/cuda/causal_conv1d.py b/lmdeploy/pytorch/kernels/cuda/causal_conv1d.py index b634310918..5b88bf0106 100644 --- a/lmdeploy/pytorch/kernels/cuda/causal_conv1d.py +++ b/lmdeploy/pytorch/kernels/cuda/causal_conv1d.py @@ -217,28 +217,33 @@ 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, + x_stride=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') 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])), - 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 +373,10 @@ 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, + 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..693b5d915c 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, @@ -53,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, @@ -69,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 @@ -93,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( @@ -138,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 @@ -154,6 +211,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, @@ -161,6 +222,7 @@ def hc_post_expand( post, comb, out, + out_fp32, *x.stride(), *residual.stride(), *post.stride(), @@ -172,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/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/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/kernels/cuda/gated_delta_rule.py b/lmdeploy/pytorch/kernels/cuda/gated_delta_rule.py index 2ec3109089..f83dd31c24 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,11 @@ 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, + 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, @@ -301,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) @@ -310,7 +319,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 +328,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,13 +466,17 @@ 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, 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 @@ -488,13 +502,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 @@ -521,8 +552,11 @@ 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: + 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 +576,17 @@ 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): + 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: + 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, @@ -585,6 +629,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. @@ -592,7 +639,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 @@ -610,6 +658,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, @@ -624,6 +676,20 @@ 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 + 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) if g is not None: assert g.is_contiguous() g_dtype = g.dtype @@ -646,7 +712,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: @@ -693,9 +760,16 @@ 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, + 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/lmdeploy/pytorch/kernels/cuda/kpool.py b/lmdeploy/pytorch/kernels/cuda/kpool.py new file mode 100644 index 0000000000..8e6a0f050d --- /dev/null +++ b/lmdeploy/pytorch/kernels/cuda/kpool.py @@ -0,0 +1,468 @@ +# 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 _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 + 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) + 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): + 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, + 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 _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, 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) + 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) + 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))) + 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, + 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. 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, 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/kernels/cuda/moe/blocked_fp8.py b/lmdeploy/pytorch/kernels/cuda/moe/blocked_fp8.py index ba0a91f5e2..641ecb8f1f 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: @@ -690,7 +695,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, + *, + output_scale: float = 1.0, + fp32_acc: bool = False) -> torch.Tensor: """Fused moe.""" device = input.device M = input.size(0) @@ -816,5 +824,7 @@ 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/ep.py b/lmdeploy/pytorch/kernels/cuda/moe/ep.py index ed51f840a8..cde008b8ee 100644 --- a/lmdeploy/pytorch/kernels/cuda/moe/ep.py +++ b/lmdeploy/pytorch/kernels/cuda/moe/ep.py @@ -142,14 +142,12 @@ 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) - # 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) accumulator = tl.zeros([BLOCK_D], dtype=compute_dtype) @@ -173,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 @@ -200,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 8eab01008f..6fd9c64723 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,22 @@ 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) + # 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), device=hidden_states_fp8.device, @@ -183,12 +189,14 @@ 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) + ep_gather(down_output, topk_idx, topk_weights, output_index, gather_out, fp32_acc=fp32_acc) return gather_out diff --git a/lmdeploy/pytorch/kernels/cuda/moe/fused_moe.py b/lmdeploy/pytorch/kernels/cuda/moe/fused_moe.py index 4626694765..c648325176 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,15 @@ 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 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) @@ -1008,6 +1013,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, ) @@ -1025,7 +1031,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, + output_scale: float = 1.0, + fp32_acc: bool = False) -> torch.Tensor: """Fused moe.""" M = hidden_states.size(0) E, N, _ = w1.shape @@ -1133,5 +1141,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/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_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/deepseek_v2.py b/lmdeploy/pytorch/models/deepseek_v2.py index ccadbfb1f5..f4727231f3 100644 --- a/lmdeploy/pytorch/models/deepseek_v2.py +++ b/lmdeploy/pytorch/models/deepseek_v2.py @@ -605,12 +605,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 @@ -704,6 +707,12 @@ 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_output_scale = 1.0 + fused_moe_fp32_acc = False + router_routed_scaling_factor = None + shared_expert_cls = None + def __init__(self, config: Any, layer_idx, @@ -736,9 +745,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, @@ -750,12 +771,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, + 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 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, diff --git a/lmdeploy/pytorch/models/deepseek_v32.py b/lmdeploy/pytorch/models/deepseek_v32.py index fd93567af0..c8dfb15fe9 100644 --- a/lmdeploy/pytorch/models/deepseek_v32.py +++ b/lmdeploy/pytorch/models/deepseek_v32.py @@ -114,7 +114,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:] @@ -249,6 +250,8 @@ def forward(self, class DeepseekV32Attention(DeepseekV2Attention): + use_sparse_mla = True + def __init__(self, config: Any, layer_idx: int, @@ -360,7 +363,8 @@ def __init__(self, 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..3c72e4ca40 --- /dev/null +++ b/lmdeploy/pytorch/models/glm5_next.py @@ -0,0 +1,2137 @@ +# 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.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, + kpool_score_contiguous_cuda, + 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 ( + 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_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, + FlashAttention, + HcPrePost, + Kda, + KPoolIndexer, + LayerNorm, + ParallelLMHead, + RMSNorm, +) +from lmdeploy.pytorch.nn.gated_delta import GatedDeltaMeta, GatedDeltaMetaBuilder, 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_batched_linear, + build_colwise_linear, + build_merged_colwise_linear, + build_o_proj, + build_qkv_proj, + build_rowwise_linear, +) +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 .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 + +Glm5NextVisionRMSNorm = RMSNorm + + +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) +_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): + """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.apply_rotary_pos_emb = ApplyRotaryEmb(enable_fp32_compute=True) + 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 = self.apply_rotary_pos_emb(query, key, cos, sin, inplace=False) + 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 = 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], + 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.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, + 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 + 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) + 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. 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 = ( + 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) + # Match the GLM-5.3 contract: routing returns normalized, unscaled weights; + # 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 + fused_moe_fp32_acc = True + shared_expert_cls = Glm5NextMLP + + 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). + 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 ' + '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.') + out = super().forward(hidden_states, all_routed_experts=None) + if self._fp32_tp_reduce: + output_dtype = out.dtype + out = out.float() + get_dist_group('moe').all_reduce_(out) + out = out.to(output_dtype) + return out + + +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, + ) + # 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], + 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.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) + 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, + ) + 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, + 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, norm_gate = self.fg_b_proj(self.fg_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 + + 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, + 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. + 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: + 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 ' + '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. 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), + bias=False, + dtype=dtype, + device=device, + is_tp=True, + quant_config=None, + ) + 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, + ) + 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, + device: torch.device, prefix: str = ''): + del layer_idx, prefix + 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, + ) + + 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, + *, + 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: + 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 + history_lengths = attn_metadata.kv_seqlens - attn_metadata.q_seqlens + indexer_k_cache = self.indexer.get_block_cache() + 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: + 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 + 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) + + if attn_metadata.is_decoding: + batch_size = state_ids.numel() + if key.size(0) % batch_size: + raise RuntimeError( + '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 + + 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): + 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)}.') + save_ring(attn_metadata.kv_seqlens) + 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, + 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 + 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 + broadcast_groups = attn_metadata.is_decoding or (hidden_states.is_cuda and dist_ctx.dist_config.attn_tp > 1) + if attn_metadata.is_decoding: + 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] + 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: + 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, + group_lengths, + group_topk=self.index_topk // self.index_kpool, + ) + logical_indices = selected_groups + else: + logical_indices = self._select_kpool_indices_prefill( + query_fp8, + 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 broadcast_groups else 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) + 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) + 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, + *, + return_groups: bool = False, + ) -> torch.Tensor: + """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, return_groups=return_groups) + 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, + 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, + 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) + 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 + indices = self._select_kpool_indices( + 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. + 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: + 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, + 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 + 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) + 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, + topk_indices_buffer=topk_indices_buffer, + skip_topk=skip_topk, + kpool_metadata=kpool_metadata, + ) + 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( + query_states, + key_states, + key_states[..., :nope_size], + past_key_value[0], + past_key_value[0][..., :nope_size], + attn_metadata, + 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) + 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, + 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) + 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, + 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) + 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, + prefix=f'model.layers.{layer_idx}.mlp')) + + 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, + 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, + 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, + x_fp32: torch.Tensor | None = None): + return self.hc_prepost.pre( + 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, + 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, + self.hc_attn_fn, + 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, + 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, + 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( + hidden_states, + self.hc_ffn_fn, + 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_post(hidden_states, residual, post, comb, prepare_hc_fp32) + + +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.gated_delta_meta_builder = GatedDeltaMetaBuilder() + 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 = self.gated_delta_meta_builder(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)}.') + 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, 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, + 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) + + def get_input_embeddings(self): + return self.embed_tokens + + +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, + 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, + return_input_embeds: bool = False, + **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)) + 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() + + 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, + 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]], *, + is_mtp: bool = False): + """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 = [ + ('.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'), + ('.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 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 + 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) + + +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]) + 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) + 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, projected=(key, score)) + + # 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')] + 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, 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, + 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 + + +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 = 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, + 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..7fc18402b2 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) }) @@ -232,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/models/module_map.py b/lmdeploy/pytorch/models/module_map.py index 74ac510934..6e6dae21b8 100644 --- a/lmdeploy/pytorch/models/module_map.py +++ b/lmdeploy/pytorch/models/module_map.py @@ -63,6 +63,13 @@ 'GlmMoeDsaForCausalLM': f'{LMDEPLOY_PYTORCH_MODEL_PATH}.glm_moe_dsa.GlmMoeDsaForCausalLM', }) +# GLM-5.3 Flash +MODULE_MAP.update({ + 'Glm5NextForConditionalGeneration': + f'{LMDEPLOY_PYTORCH_MODEL_PATH}.glm5_next.Glm5NextForConditionalGeneration', + 'Glm5NextMTPModel': f'{LMDEPLOY_PYTORCH_MODEL_PATH}.glm5_next.Glm5NextMTPModel', +}) + # internlm2 MODULE_MAP.update({ 'InternLM2ForCausalLM': f'{LMDEPLOY_PYTORCH_MODEL_PATH}.internlm2.InternLM2ForCausalLM', diff --git a/lmdeploy/pytorch/nn/__init__.py b/lmdeploy/pytorch/nn/__init__.py index c167e588ec..1feb2c55e2 100644 --- a/lmdeploy/pytorch/nn/__init__.py +++ b/lmdeploy/pytorch/nn/__init__.py @@ -5,6 +5,8 @@ from .attention import Attention, FlashAttention # noqa: F401 from .embedding import ParallelEmbedding, ParallelLMHead # noqa: F401 from .hc_prepost import HcPrePost # noqa: F401 +from .kda import Kda # noqa: F401 +from .kpool import KPoolIndexer # noqa: F401 from .norm import LayerNorm, RMSNorm, rms_scale # noqa: F401 from .rotary_embedding import ( ApplyRotaryEmb, # noqa: F401 diff --git a/lmdeploy/pytorch/nn/attention.py b/lmdeploy/pytorch/nn/attention.py index 494982875b..0d89e25a7e 100644 --- a/lmdeploy/pytorch/nn/attention.py +++ b/lmdeploy/pytorch/nn/attention.py @@ -93,6 +93,37 @@ 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/hc_prepost.py b/lmdeploy/pytorch/nn/hc_prepost.py index 97f0c717ab..f6feaf6ecb 100644 --- a/lmdeploy/pytorch/nn/hc_prepost.py +++ b/lmdeploy/pytorch/nn/hc_prepost.py @@ -8,11 +8,18 @@ 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.""" - 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, @@ -25,12 +32,28 @@ def pre( hc_scale: torch.Tensor, 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 - shape, dtype = x.size(), x.dtype - x = x.flatten(2).float() - mixes = rms_scale(F.linear(x, hc_fn), x, eps=norm_eps) - return self.impl.pre(x.view(shape), mixes, hc_scale, hc_base, dtype) + 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. + 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. + 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(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) @@ -38,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/kda.py b/lmdeploy/pytorch/nn/kda.py new file mode 100644 index 0000000000..17072d109e --- /dev/null +++ b/lmdeploy/pytorch/nn/kda.py @@ -0,0 +1,52 @@ +# Copyright (c) OpenMMLab. All rights reserved. +from typing import Any + +import torch +from torch import nn + +from lmdeploy.pytorch.backends import get_backend +from lmdeploy.pytorch.backends.kda import KdaBuildSpec +from lmdeploy.pytorch.models.patch import get_build_model_context + + +class Kda(nn.Module): + """Backend-dispatched Kimi Delta Attention recurrence.""" + + def __init__(self): + super().__init__() + self.impl = get_backend().build_op( + KdaBuildSpec(), + enable_deterministic=get_build_model_context().enable_deterministic, + ) + + 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..80bd718766 --- /dev/null +++ b/lmdeploy/pytorch/nn/kpool.py @@ -0,0 +1,925 @@ +# 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, DSA_INDEXER_K_CACHE_NAME, dsa_packed_indexer_k_cache_shape +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 + +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 + self._block_cache_binding: BlockCacheBinding | None = None + + 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 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, + ) + self.index_kpool_compress_gate = nn.Parameter( + torch.empty(index_head_dim, hidden_size, dtype=dtype, device=device), + requires_grad=False, + ) + + def get_block_cache_requests(self, context: BlockCacheRequestContext): + """Declare the pooled index cache through the shared cache planner.""" + geometry = context.geometry + if geometry.logical_block_size != 64 or geometry.kernel_block_size != 64: + raise ValueError('GLM-5.3 KPool requires logical and kernel block_size=64.') + return (BlockCacheRequest( + name=DSA_INDEXER_K_CACHE_NAME, + shape=dsa_packed_indexer_k_cache_shape(64, self.head_dim), + dtype=torch.uint8, + per_row_contiguous=True, + ), ) + + def bind_block_cache(self, binding: BlockCacheBinding): + """Retain the compact consumer row assigned by the cache planner.""" + if binding.cache_name != DSA_INDEXER_K_CACHE_NAME: + raise ValueError(f'Unexpected KPool cache name: {binding.cache_name}.') + self._block_cache_binding = binding + + def get_block_cache(self) -> Tensor: + """Resolve this indexer's row from the live request context.""" + binding = self._block_cache_binding + if binding is None: + raise RuntimeError('The KPool index cache has not been bound.') + caches = get_step_ctx_manager().current_context().block_caches + if hasattr(caches, 'row'): + return caches.row(binding.cache_name, binding.consumer_row) + return caches[binding.cache_name][binding.consumer_row] + + 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.""" + 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.""" + 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/__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/base.py b/lmdeploy/pytorch/nn/linear/base.py index 03ec93bb0f..6dbd7b3ed9 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,13 @@ 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: - dist.all_reduce(out, group=self.tp_group) + output_dtype = out.dtype + if self.tp_reduce_dtype is not None: + out = out.to(self.tp_reduce_dtype) + 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 def _forward_dp_tp(self, x): @@ -232,7 +242,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/nn/linear/default.py b/lmdeploy/pytorch/nn/linear/default.py index efa3ea3bea..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.""" @@ -128,17 +163,20 @@ def get_unquantized_weight(self, out_dtype: torch.dtype) -> torch.Tensor: 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..7214dabece 100644 --- a/lmdeploy/pytorch/nn/moe/__init__.py +++ b/lmdeploy/pytorch/nn/moe/__init__.py @@ -24,6 +24,9 @@ def build_fused_moe( layer_idx: int = 0, act_func: Callable = None, prefix: str = '', + *, + output_scale: float = 1.0, + fp32_acc: bool = False, ): """Fused moe builder.""" quant_method = None @@ -45,9 +48,13 @@ def build_fused_moe( all_reduce=all_reduce, layer_idx=layer_idx, act_func=act_func, + output_scale=output_scale, + fp32_acc=fp32_acc, ) 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 @@ -69,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.' ) @@ -107,8 +116,13 @@ def build_fused_moe( all_reduce=all_reduce, 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: 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 94d0bb6208..da292b0513 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, @@ -157,7 +159,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, + output_scale: float = 1.0, + fp32_acc: bool = False): device = device or torch.device('cpu') dtype = dtype or torch.float16 @@ -190,6 +194,8 @@ def __init__(self, layer_idx=layer_idx, 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, ) @@ -230,6 +236,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): @@ -252,6 +259,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 @@ -334,7 +344,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'], @@ -345,7 +356,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/nn/moe/default.py b/lmdeploy/pytorch/nn/moe/default.py index 7705a0f082..02d7da928d 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, + output_scale: float = 1.0, + fp32_acc: bool = False): 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, + output_scale=output_scale, + fp32_acc=fp32_acc, ), enable_deterministic=build_ctx.enable_deterministic, ) diff --git a/lmdeploy/pytorch/nn/rotary_embedding.py b/lmdeploy/pytorch/nn/rotary_embedding.py index 466c70fbac..2757e2c085 100644 --- a/lmdeploy/pytorch/nn/rotary_embedding.py +++ b/lmdeploy/pytorch/nn/rotary_embedding.py @@ -245,12 +245,13 @@ def build_rotary_embedding_from_config(config: PretrainedConfig, device: torch.d 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/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, 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: 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 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..7effc460c9 --- /dev/null +++ b/lmdeploy/vl/model/glm5_next.py @@ -0,0 +1,286 @@ +# 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'] + + @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 a1c56069ac..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>=0.4.2 +flash-linear-attention==0.5.2 opencv-python-headless peft<=0.14.0 prometheus_client