Skip to content
Closed
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
172 changes: 144 additions & 28 deletions vllm/models/glm5next/nvidia/kda.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,7 @@
from vllm.third_party.flash_linear_attention.ops.kda import FusedRMSNormGated
from vllm.transformers_utils.configs.glm5_next import Glm5NextConfig
from vllm.v1.attention.backends.gdn_attn import GDNAttentionMetadata
from vllm.v1.worker.workspace import current_workspace_manager

if current_platform.is_rocm():
from vllm.models.glm5next.amd.ops.third_party.kda import (
Expand Down Expand Up @@ -123,6 +124,32 @@ def _cast_sigmoid(x: torch.Tensor) -> torch.Tensor:
return x.float().sigmoid()


def _resolve_kda_prefill_backend(
backend: str, head_dim: int, dtype: torch.dtype, lower_bound: float | None
) -> str:
"""Pick the chunked-prefill kernel: FlashKDA (fused CUDA, ~2-4x faster on
SM90/SM10x/SM12x for bf16, head_dim 128 and a bounded gate) or the Triton
``chunk_kda_with_fused_gate`` path. ``backend`` comes from
``additional_config.kda_prefill_backend`` (auto / triton / flashkda)."""
if backend not in ("auto", "triton", "flashkda"):
raise ValueError(f"Unsupported KDA prefill backend: {backend}")
capability = current_platform.get_device_capability()
supported = (
current_platform.is_cuda()
and capability is not None
and capability.major in (9, 10, 12)
and head_dim == 128
and dtype == torch.bfloat16
and lower_bound is not None
)
if backend == "flashkda" and not supported:
raise RuntimeError(
"FlashKDA requires CUDA SM90/SM10x/SM12x, bfloat16, head_dim=128 "
"and a bounded KDA gate."
)
return "flashkda" if supported and backend != "triton" else "triton"
Comment on lines +127 to +150


class Glm5NextLinearAttention(GatedDeltaNetAttention):
head_dim: int
num_heads: int
Expand Down Expand Up @@ -282,6 +309,85 @@ def _a_log_weight_loader(param, loaded_weight):
# every _forward call (it reads an env-derived flag each time).
self._conv_state_dim_first = is_conv_state_dim_first()

additional_config = vllm_config.additional_config
self.kda_prefill_backend = _resolve_kda_prefill_backend(
additional_config.get("kda_prefill_backend", "auto")
if isinstance(additional_config, dict)
else "auto",
self.head_dim,
vllm_config.model_config.dtype,
self.kda_lower_bound,
)
self._flashkda_buffer_specs: (
tuple[tuple[tuple[int, ...], torch.dtype], ...] | None
) = None
if self.kda_prefill_backend == "flashkda":
import vllm._flashkda_C # noqa: F401

max_tokens = vllm_config.scheduler_config.max_num_batched_tokens
max_seqs = vllm_config.scheduler_config.max_num_seqs
workspace_size = torch.ops._flashkda_C.get_workspace_size(
max_tokens, self.local_num_heads, max_seqs
)
self._flashkda_buffer_specs = (
(
(max_seqs, self.local_num_heads, self.head_dim, self.head_dim),
self.get_state_dtype()[1],
),
((workspace_size,), torch.uint8),
# Output buffer for steps that also carry spec-decode tokens:
# the non-spec tokens are then scattered by non_spec_token_indx.
(
(1, max_tokens, self.local_num_heads, self.head_dim),
vllm_config.model_config.dtype,
),
)

def _flashkda_prefill(
self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
initial_state: torch.Tensor,
cu_seqlens: torch.Tensor,
out: torch.Tensor | None,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Fused KDA chunked prefill (FlashKDA). Takes the raw gate logits ``g``
and raw ``beta`` logits, l2-normalizes q/k in-kernel and applies the
bounded gate ``lower_bound * sigmoid(exp(A_log) * (g + dt_bias))``,
matching ``chunk_kda_with_fused_gate(..., safe_gate=True)``. Writes the
attention output into ``out`` (a workspace buffer when ``None``) and
returns ``(out, final_state)``."""
assert self._flashkda_buffer_specs is not None
final_state, workspace, workspace_out = (
current_workspace_manager().get_simultaneous(*self._flashkda_buffer_specs)
)
final_state = final_state[: initial_state.shape[0]]
if out is None:
out = workspace_out[:, : q.shape[1]]
# FlashKDA hardcodes dense q/k/v/g strides; beta may be row-strided.
torch.ops._flashkda_C.fwd(
q.contiguous(),
k.contiguous(),
v.contiguous(),
g.contiguous(),
beta,
self.head_dim**-0.5,
out,
workspace,
self.A_log.view(-1),
self.dt_bias.view(-1, self.head_dim),
self.kda_lower_bound,
initial_state.contiguous(),
final_state,
cu_seqlens.contiguous(),
None,
None,
)
return out, final_state

def forward(
self,
hidden_states: torch.Tensor,
Expand Down Expand Up @@ -533,26 +639,42 @@ def _rearr(x):
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,
)
if self.kda_prefill_backend == "flashkda":
# Non-spec step: write straight into the layer output buffer
# (dense token order, no merge copy). Step with spec-decode
# tokens: write to the workspace buffer and scatter below.
ns_out = None if use_spec else core_attn_out[:, :num_actual_tokens]
core_attn_out_non_spec, last_recurrent_state = self._flashkda_prefill(
q=_rearr(q_ns),
k=_rearr(k_ns),
v=_rearr(v_ns),
g=g1_ns,
beta=beta_ns,
initial_state=initial_state,
cu_seqlens=non_spec_query_start_loc,
out=ns_out,
)
else:
(
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,
Expand Down Expand Up @@ -591,14 +713,8 @@ def _rearr(x):
# --- 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)
core_attn_out.index_copy_(1, spec_token_indx, core_attn_out_spec)
core_attn_out.index_copy_(1, non_spec_token_indx, core_attn_out_non_spec)
elif use_spec:
assert core_attn_out_spec is not None
if spec_out is None:
Expand Down