# SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project """GLM5-Next KDA (linear-attention) layer. Model-specific, self-contained KDA: separate q/k/v short-conv + the GLM5-Next spec-decode verify path + the bounded ``safe_gate`` variant, and ``_forward`` is an eager break point under Breakable CUDA Graph (``@eager_break_during_capture``). Moved out of the shared ``kimi_gdn_linear_attn.py`` (which reverts to Kimi Linear's fused-conv version): the separate-conv layout + spec-verify are GLM5-Next-only. ``forward`` calls ``self._forward`` directly (no ``torch.ops.vllm.kda_attention`` indirection) so the only un-capturable work is the decorated ``_forward``. """ import torch from torch import nn from vllm.compilation.breakable_cudagraph import eager_break_during_capture from vllm.config import VllmConfig, get_current_vllm_config from vllm.distributed import divide from vllm.forward_context import get_forward_context from vllm.model_executor.layers.linear import ( ColumnParallelLinear, MergedColumnParallelLinear, RowParallelLinear, ) from vllm.model_executor.layers.mamba.gdn.base import GatedDeltaNetAttention from vllm.model_executor.layers.mamba.mamba_utils import ( MambaStateDtypeCalculator, MambaStateShapeCalculator, is_conv_state_dim_first, ) from vllm.model_executor.layers.mamba.ops.causal_conv1d import ( causal_conv1d_fn, causal_conv1d_update, ) from vllm.model_executor.layers.mamba.ops.gather_initial_states import ( gather_initial_states, ) from vllm.model_executor.layers.mamba.ops.scatter_states import scatter_states from vllm.model_executor.model_loader.weight_utils import sharded_weight_loader from vllm.model_executor.utils import ( maybe_disable_graph_partition, set_weight_attrs, ) from vllm.platforms import current_platform from vllm.third_party.flash_linear_attention.ops.kda import ( FusedRMSNormGated, chunk_kda_with_fused_gate, fused_recurrent_kda, ) from vllm.transformers_utils.configs.kimi_linear import KimiLinearConfig from vllm.v1.attention.backends.gdn_attn import GDNAttentionMetadata class _Glm5NextMergedColumnParallelLinear(MergedColumnParallelLinear): """Merged projection with multiple replicated output shards. Extends K3's ``_KimiGDNMergedColumnParallelLinear`` to support two replicated shards (f_a, g_a) instead of one. Pre-multiplies each replicated entry's output_size by tp_size so the per-rank shard divides back to the full size, and forces tp_rank=0 during weight loading for replicated shards. """ def __init__( self, input_size: int, output_sizes: list[int], replicated_shard_ids: tuple[int, ...], tp_size: int, **kwargs, ) -> None: self.replicated_shard_ids = set(replicated_shard_ids) output_sizes = output_sizes.copy() for sid in self.replicated_shard_ids: output_sizes[sid] *= tp_size super().__init__(input_size, output_sizes, **kwargs) def weight_loader( self, param: nn.Parameter, loaded_weight: torch.Tensor, loaded_shard_id: tuple[int, ...] | int | None = None, ) -> None: tp_rank = self.tp_rank param_tp_rank = getattr(param, "tp_rank", None) if loaded_shard_id in self.replicated_shard_ids: self.tp_rank = 0 if param_tp_rank is not None: param.tp_rank = 0 try: super().weight_loader(param, loaded_weight, loaded_shard_id) finally: self.tp_rank = tp_rank if param_tp_rank is not None: param.tp_rank = param_tp_rank def weight_loader_v2( self, param: nn.Parameter, loaded_weight: torch.Tensor, loaded_shard_id: tuple[int, ...] | int | None = None, ) -> None: tp_rank = self.tp_rank param_tp_rank = getattr(param, "tp_rank", None) if loaded_shard_id in self.replicated_shard_ids: self.tp_rank = 0 if param_tp_rank is not None: param.tp_rank = 0 try: super().weight_loader_v2(param, loaded_weight, loaded_shard_id) finally: self.tp_rank = tp_rank if param_tp_rank is not None: param.tp_rank = param_tp_rank @torch.compile( dynamic=True, backend=current_platform.simple_compile_backend, options=maybe_disable_graph_partition(current_platform.simple_compile_backend), ) def _cast_sigmoid(x: torch.Tensor) -> torch.Tensor: """Fuse the fp32 cast + sigmoid into one Inductor kernel.""" return x.float().sigmoid() class Glm5NextLinearAttention(GatedDeltaNetAttention): # Declared int (set in __init__ from config) so mypy doesn't see the # getattr-derived `Any | None` at the kernel call sites. head_dim: int num_heads: int conv_size: int def get_state_dtype( self, ) -> tuple[torch.dtype, torch.dtype]: if self.model_config is None or self.cache_config is None: raise ValueError("model_config and cache_config must be set") return MambaStateDtypeCalculator.kda_state_dtype( self.model_config.dtype, self.cache_config.mamba_cache_dtype ) def get_state_shape( self, ) -> tuple[tuple[int, ...], tuple[int, ...]]: # conv_state width must include num_spec so the spec-decode conv update # (causal_conv1d_update with num_accepted_tokens + max_query_len) can # slide the window across the draft-verify tokens without reading past # the allocated width. Matches qwen_gdn_linear_attn.get_state_shape. return MambaStateShapeCalculator.kda_state_shape( self.tp_size, self.num_heads, self.head_dim, conv_kernel_size=self.conv_size, num_spec=self.num_spec, ) def __init__( self, config: KimiLinearConfig, vllm_config: VllmConfig, prefix: str = "", ) -> None: # LOCAL PATCH (fp8attn-r2): pass the real quant config through so KDA # projections can be FP8-resident when the checkpoint declares them in # a MIXED_PRECISION quantized_layers manifest. Checkpoints that keep # attention BF16 (e.g. the stock NVFP4 export) are unaffected: their # quantization_config `ignore` list names every self_attn module # (including the fused `in_proj_qkvbfg_a` spelling), so every KDA # linear still resolves to UnquantizedLinearMethod. # Was: save/None/restore strip of vllm_config.quant_config. super().__init__(config, vllm_config, prefix) # Linear-attention head config: read the flattened top-level fields when # present (new schema); fall back to the legacy linear_attn_config dict # otherwise (shared base is also used by KimiLinearConfig). Narrow via # locals so the int-typed attrs are assigned a non-None value. head_dim = getattr(config, "linear_head_dim", None) num_heads = getattr(config, "linear_num_heads", None) conv_size = getattr(config, "linear_conv_kernel_dim", None) if head_dim is None or num_heads is None or conv_size is None: kda_config = config.linear_attn_config # type: ignore[attr-defined] assert kda_config is not None, "linear_attn_config must be set" head_dim = kda_config["head_dim"] num_heads = kda_config["num_heads"] conv_size = kda_config["short_conv_kernel_size"] assert head_dim is not None assert num_heads is not None assert conv_size is not None self.head_dim = head_dim self.num_heads = num_heads self.conv_size = conv_size assert self.num_heads % self.tp_size == 0 self.local_num_heads = divide(self.num_heads, self.tp_size) projection_size = self.head_dim * self.num_heads self.local_projection_size = divide(projection_size, self.tp_size) # Merge q, k, v, b, f_a, g_a projections into one GEMM (6→1 launches). # Order matches checkpoint's fused_qkvbfg_a_proj convention. # Shards 4 (f_a) and 5 (g_a) are replicated across TP ranks. self.in_proj_qkvbfg_a = _Glm5NextMergedColumnParallelLinear( self.hidden_size, [ projection_size, # q (shard 0) projection_size, # k (shard 1) projection_size, # v (shard 2) self.num_heads, # b (shard 3) self.head_dim, # f_a (shard 4, replicated) self.head_dim, # g_a (shard 5, replicated) ], replicated_shard_ids=(4, 5), tp_size=self.tp_size, bias=False, quant_config=self.quant_config, prefix=f"{prefix}.in_proj_qkvbfg_a", ) self.f_b_proj = ColumnParallelLinear( self.head_dim, projection_size, bias=False, quant_config=self.quant_config, prefix=f"{prefix}.f_b_proj", ) self.dt_bias = nn.Parameter( torch.empty(divide(projection_size, self.tp_size), dtype=torch.float32) ) set_weight_attrs(self.dt_bias, {"weight_loader": sharded_weight_loader(0)}) self.q_conv1d = ColumnParallelLinear( input_size=self.conv_size, output_size=projection_size, bias=False, params_dtype=torch.float32, prefix=f"{prefix}.q_conv1d", ) self.k_conv1d = ColumnParallelLinear( input_size=self.conv_size, output_size=projection_size, bias=False, params_dtype=torch.float32, prefix=f"{prefix}.k_conv1d", ) self.v_conv1d = ColumnParallelLinear( input_size=self.conv_size, output_size=projection_size, bias=False, params_dtype=torch.float32, prefix=f"{prefix}.v_conv1d", ) # unsqueeze to fit conv1d weights shape into the linear weights shape. # Can't do this in `weight_loader` since it already exists in # `ColumnParallelLinear` and `set_weight_attrs` # doesn't allow to override it self.q_conv1d.weight.data = self.q_conv1d.weight.data.unsqueeze(1) self.k_conv1d.weight.data = self.k_conv1d.weight.data.unsqueeze(1) self.v_conv1d.weight.data = self.v_conv1d.weight.data.unsqueeze(1) # Lazily-built merged q|k|v conv weight (built on first forward, after # weights are loaded). See _forward. self._merged_conv_weight: torch.Tensor | None = None self.A_log = nn.Parameter( torch.empty(1, 1, self.local_num_heads, 1, dtype=torch.float32) ) set_weight_attrs(self.A_log, {"weight_loader": sharded_weight_loader(2)}) self.g_b_proj = ColumnParallelLinear( self.head_dim, projection_size, bias=False, quant_config=self.quant_config, prefix=f"{prefix}.g_b_proj", ) self.o_norm = FusedRMSNormGated(self.head_dim, activation="sigmoid") self.o_proj = RowParallelLinear( projection_size, self.hidden_size, bias=False, quant_config=self.quant_config, prefix=f"{prefix}.o_proj", ) compilation_config = get_current_vllm_config().compilation_config if prefix in compilation_config.static_forward_context: raise ValueError(f"Duplicate layer name: {prefix}") compilation_config.static_forward_context[prefix] = self # GLM5-Next checkpoints A_log as 1-D (num_heads,); the param is 4-D, so # reshape on load before the sharded loader runs. def _a_log_weight_loader(param, loaded_weight): if loaded_weight.dim() == 1: loaded_weight = loaded_weight.view([1, 1, -1, 1]) return sharded_weight_loader(2)(param, loaded_weight) self.A_log.weight_loader = _a_log_weight_loader # Bounded KDA gate variant: GLM5-Next uses # y = lower_bound * sigmoid(exp(A)*(g+g_bias)) instead of the default # unbounded y = -exp(A)*softplus(g+g_bias). Read by _forward. linear_lower_bound = getattr(config, "linear_lower_bound", None) if linear_lower_bound is not None: self.kda_safe_gate = True self.kda_lower_bound = linear_lower_bound else: legacy = getattr(config, "linear_attn_config", None) or {} if legacy.get("safe_gate", True): self.kda_safe_gate = True self.kda_lower_bound = legacy.get("lower_bound", -5.0) else: self.kda_safe_gate = False self.kda_lower_bound = -5.0 # Process-global conv-state layout, resolved once here instead of on # every _forward call (it reads an env-derived flag each time). self._conv_state_dim_first = is_conv_state_dim_first() def forward( self, hidden_states: torch.Tensor, positions: torch.Tensor, ) -> torch.Tensor: num_tokens = hidden_states.size(0) # One merged GEMM for q, k, v, b, f_a, g_a (replaces 6 separate GEMMs). projected = self.in_proj_qkvbfg_a(hidden_states)[0] qkv, beta_raw, f_a, g_a = projected.split( [ 3 * self.local_projection_size, self.local_num_heads, self.head_dim, self.head_dim, ], dim=-1, ) # Beta stays raw (bf16) here: the recurrent kernel sigmoids it in fp32 # at load (SIGMOID_BETA), and only the chunked prefill path needs the # pre-computed fp32 sigmoid — computed lazily in _forward. Pure decode # / spec-verify steps then skip the _cast_sigmoid kernel and its fp32 # intermediate entirely. beta = beta_raw.unsqueeze(0) g1 = self.f_b_proj(f_a)[0] g1 = g1.reshape(1, -1, self.local_num_heads, self.head_dim) g_proj_states = self.g_b_proj(g_a)[0] # Must stay 3D: rms_norm_gated reads H from g.shape[-2]. g2 = g_proj_states.reshape(-1, self.local_num_heads, self.head_dim) core_attn_out = torch.empty( (1, num_tokens, self.local_num_heads, self.head_dim), dtype=hidden_states.dtype, device=hidden_states.device, ) # Call _forward directly (not via the registered op) so the KDA core # is an eager break point under Breakable CG, mirroring KimiK3's KDA # (vllm/models/kimi_k3/nvidia/kda.py). torch.ops.vllm.kda_attention is # neither a splitting op nor @eager_break_during_capture-decorated, so # routing through it lets the host-branching prefill body be # Inductor-compiled + stream-captured under PIECEWISE -> stale garbage. # qkv stays merged through the short-conv (one conv call, not three). self._forward( qkv_proj_states=qkv, g1=g1, beta=beta, core_attn_out=core_attn_out, ) core_attn_out = self.o_norm(core_attn_out, g2) core_attn_out = core_attn_out.reshape(core_attn_out.size(1), -1) return self.o_proj(core_attn_out)[0] @eager_break_during_capture def _forward( self, qkv_proj_states: torch.Tensor, g1: torch.Tensor, beta: torch.Tensor, core_attn_out: torch.Tensor, ) -> None: forward_context = get_forward_context() attn_metadata_raw = forward_context.attn_metadata if attn_metadata_raw is None: # # V1 profile run return assert isinstance(attn_metadata_raw, dict) attn_metadata_narrowed = attn_metadata_raw[self.prefix] assert isinstance(attn_metadata_narrowed, GDNAttentionMetadata) has_initial_state = attn_metadata_narrowed.has_initial_state non_spec_query_start_loc = attn_metadata_narrowed.non_spec_query_start_loc non_spec_state_indices_tensor = ( attn_metadata_narrowed.non_spec_state_indices_tensor ) # noqa: E501 num_actual_tokens = attn_metadata_narrowed.num_actual_tokens # Spec-decode metadata (all None when speculative decoding is disabled). spec_sequence_masks = attn_metadata_narrowed.spec_sequence_masks spec_query_start_loc = attn_metadata_narrowed.spec_query_start_loc spec_state_indices_tensor = attn_metadata_narrowed.spec_state_indices_tensor spec_token_indx = attn_metadata_narrowed.spec_token_indx non_spec_token_indx = attn_metadata_narrowed.non_spec_token_indx num_accepted_tokens = attn_metadata_narrowed.num_accepted_tokens num_spec_decodes = attn_metadata_narrowed.num_spec_decodes use_spec = spec_sequence_masks is not None and num_spec_decodes > 0 # KDA gate variant: GLM5-Next checkpoints with # linear_attn_config["safe_gate"]=True use the bounded gate # y=lower_bound*sigmoid(exp(A)*(g+g_bias)) instead of the default # unbounded y=-exp(A)*softplus(g+g_bias). Both attrs are always set # in __init__ (this class is GLM5Next-only). safe_gate = self.kda_safe_gate lower_bound = self.kda_lower_bound constant_caches = self.kv_cache qkv_proj_states = qkv_proj_states[:num_actual_tokens] g1 = g1[:, :num_actual_tokens] beta = beta[:, :num_actual_tokens] (conv_state, recurrent_state) = constant_caches # conv_state must be (..., dim, width-1) for the conv kernels. # DS layout stores it that way directly; SD layout needs a transpose. # Layout is process-global and resolved once at init (see __init__). if not self._conv_state_dim_first: conv_state = conv_state.transpose(-1, -2) # One merged short-conv over q|k|v instead of three separate calls. The # 1D conv is independent per channel, so concatenating q/k/v along the # channel dim and running a single causal_conv1d is bit-identical to # three calls. The merged weight is q|k|v conv weights concatenated; # built once and cached (params are fixed after load). conv_state is # already stored as the merged q|k|v state, so it is used directly. if self._merged_conv_weight is None: def _w(m): return m.weight.view(m.weight.size(0), m.weight.size(2)) self._merged_conv_weight = torch.cat( [_w(self.q_conv1d), _w(self.k_conv1d), _w(self.v_conv1d)], dim=0, ).contiguous() conv_weights = self._merged_conv_weight conv_bias = self.q_conv1d.bias # Split projections / gating into spec (draft-verify) and non-spec token # groups when speculative decoding is active. Spec tokens carry # num_spec+1 recurrent-state columns each and are advanced with # num_accepted_tokens for rejection-sampling rollback; non-spec tokens # are one-per-request. Mirrors olmo_gdn_linear_attn.py. Projections are # [n, *] (token dim 0); g1/beta are [1, n, h, d] (token dim 1). if use_spec: # In a pure spec-verify step (no non-spec tokens) the metadata # builder sets spec_token_indx = arange(num_actual_tokens), making # the index_select calls below identity copies. Skip them on this # steady-state decode hot path. The outputs alias the inputs here; # the downstream conv/recurrent kernels read them without mutating # in place, so the aliasing is safe. if non_spec_token_indx is None or non_spec_token_indx.numel() == 0: qkv_spec = qkv_proj_states g1_spec = g1 beta_spec = beta else: qkv_spec = qkv_proj_states.index_select(0, spec_token_indx) g1_spec = g1.index_select(1, spec_token_indx) beta_spec = beta.index_select(1, spec_token_indx) if non_spec_token_indx is not None and non_spec_token_indx.numel() > 0: qkv_ns = qkv_proj_states.index_select(0, non_spec_token_indx) g1_ns = g1.index_select(1, non_spec_token_indx) beta_ns = beta.index_select(1, non_spec_token_indx) else: qkv_ns = g1_ns = beta_ns = None else: qkv_spec = g1_spec = beta_spec = None qkv_ns, g1_ns, beta_ns = qkv_proj_states, g1, beta # --- causal conv1d: spec (draft-verify) path --- if use_spec: assert spec_state_indices_tensor is not None assert num_accepted_tokens is not None conv_idx = spec_state_indices_tensor[:, 0][:num_spec_decodes] conv_mql = spec_state_indices_tensor.size(-1) qkv_spec = causal_conv1d_update( qkv_spec, conv_state, conv_weights, conv_bias, activation="silu", conv_state_indices=conv_idx, num_accepted_tokens=num_accepted_tokens, query_start_loc=spec_query_start_loc, max_query_len=conv_mql, ) q_spec, k_spec, v_spec = qkv_spec.split(self.local_projection_size, dim=-1) # --- causal conv1d: non-spec path (prefill or plain decode) --- q_ns = k_ns = v_ns = None if attn_metadata_narrowed.num_prefills > 0: assert qkv_ns is not None qkv_ns = causal_conv1d_fn( qkv_ns.transpose(0, 1), conv_weights, conv_bias, activation="silu", conv_states=conv_state, has_initial_state=has_initial_state, cache_indices=non_spec_state_indices_tensor, query_start_loc=non_spec_query_start_loc, metadata=attn_metadata_narrowed, ).transpose(0, 1) q_ns, k_ns, v_ns = qkv_ns.split(self.local_projection_size, dim=-1) elif attn_metadata_narrowed.num_decodes > 0: assert non_spec_state_indices_tensor is not None decode_conv_indices = non_spec_state_indices_tensor[ : attn_metadata_narrowed.num_decodes ] qkv_ns = causal_conv1d_update( qkv_ns, conv_state, conv_weights, conv_bias, activation="silu", conv_state_indices=decode_conv_indices, ) q_ns, k_ns, v_ns = qkv_ns.split(self.local_projection_size, dim=-1) def _rearr(x): return x.reshape(1, -1, self.local_num_heads, self.head_dim) # --- core attention: spec (draft-verify) path --- core_attn_out_spec = None # In a pure spec-verify step (no non-spec tokens) the recurrent kernel # can write straight into the layer output buffer, skipping the # fresh allocation + copy below. Mixed steps must scatter via # spec_token_indx, so they keep the kernel-managed output. spec_out = ( core_attn_out[0, :num_actual_tokens].unsqueeze(0) if non_spec_token_indx is None or non_spec_token_indx.numel() == 0 else None ) if use_spec: assert spec_state_indices_tensor is not None assert num_accepted_tokens is not None assert spec_query_start_loc is not None # Gate computed inside the recurrent kernel (COMPUTE_GATE) from # raw g1 — replicates fused_kda_gate's arithmetic bit-for-bit and # skips its launch + fp32 [n, H, D] intermediate per layer. core_attn_out_spec, _ = fused_recurrent_kda( q=_rearr(q_spec), k=_rearr(k_spec), v=_rearr(v_spec), g=g1_spec, beta=beta_spec, initial_state=recurrent_state, use_qk_l2norm_in_kernel=True, cu_seqlens=spec_query_start_loc[: num_spec_decodes + 1], ssm_state_indices=spec_state_indices_tensor, num_accepted_tokens=num_accepted_tokens, out=spec_out, sigmoid_beta=True, a_log=self.A_log, g_bias=self.dt_bias, compute_gate=True, lower_bound=lower_bound, ) # --- core attention: non-spec path (prefill or plain decode) --- core_attn_out_non_spec = None # Only the plain-decode recurrent kernel can write straight into the # layer output buffer; the chunked prefill kernel cannot, so this # stays None there and the merge copy below runs as before. ns_out = None if attn_metadata_narrowed.num_prefills > 0: assert q_ns is not None assert non_spec_state_indices_tensor is not None assert has_initial_state is not None initial_state = gather_initial_states( recurrent_state, non_spec_state_indices_tensor, has_initial_state ) ( core_attn_out_non_spec, last_recurrent_state, ) = chunk_kda_with_fused_gate( q=_rearr(q_ns), k=_rearr(k_ns), v=_rearr(v_ns), raw_g=g1_ns, # Chunk path wants the pre-sigmoided fp32 beta (its kernels # don't sigmoid); beta_ns is raw bf16 from forward. beta=_cast_sigmoid(beta_ns.squeeze(0)).unsqueeze(0), A_log=self.A_log, g_bias=self.dt_bias, initial_state=initial_state, output_final_state=True, use_qk_l2norm_in_kernel=True, cu_seqlens=non_spec_query_start_loc, safe_gate=safe_gate, lower_bound=lower_bound, ) # Init cache scatter_states( recurrent_state, last_recurrent_state, non_spec_state_indices_tensor, ) elif attn_metadata_narrowed.num_decodes > 0: assert non_spec_query_start_loc is not None assert non_spec_state_indices_tensor is not None # Plain decode step (no spec tokens): token order is dense, so the # kernel can write straight into the layer output buffer. A mixed # step scatters non-spec output via non_spec_token_indx instead. # Gate computed in-kernel (COMPUTE_GATE), beta sigmoided in-kernel. if not use_spec: ns_out = spec_out core_attn_out_non_spec, _ = fused_recurrent_kda( q=_rearr(q_ns), k=_rearr(k_ns), v=_rearr(v_ns), g=g1_ns, beta=beta_ns, initial_state=recurrent_state, use_qk_l2norm_in_kernel=True, cu_seqlens=non_spec_query_start_loc[ : attn_metadata_narrowed.num_decodes + 1 ], ssm_state_indices=non_spec_state_indices_tensor, out=ns_out, sigmoid_beta=True, a_log=self.A_log, g_bias=self.dt_bias, compute_gate=True, lower_bound=lower_bound, ) # --- merge spec / non-spec outputs back into token order --- if use_spec and core_attn_out_non_spec is not None: assert core_attn_out_spec is not None merged = torch.empty( (1, num_actual_tokens, *core_attn_out_spec.shape[2:]), dtype=core_attn_out_non_spec.dtype, device=core_attn_out_non_spec.device, ) merged.index_copy_(1, spec_token_indx, core_attn_out_spec) merged.index_copy_(1, non_spec_token_indx, core_attn_out_non_spec) core_attn_out[0, :num_actual_tokens] = merged.squeeze(0) elif use_spec: assert core_attn_out_spec is not None if spec_out is None: core_attn_out[0, :num_actual_tokens] = core_attn_out_spec.squeeze(0) else: assert core_attn_out_non_spec is not None if ns_out is None: core_attn_out[0, :num_actual_tokens] = core_attn_out_non_spec[ 0, :num_actual_tokens ]