init v0.23.0

Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
This commit is contained in:
2026-08-27 15:11:51 +08:00
parent b582a8e7d1
commit 7f8a1b1f7a
2849 changed files with 712887 additions and 22001 deletions

View File

@@ -0,0 +1,11 @@
from .chunk_gated_delta_rule import chunk_gated_delta_rule_310
from .fused_gdn_gating import fused_gdn_gating_pytorch
from .fused_recurrent_gated_delta_rule import fused_recurrent_gated_delta_rule_pytorch
from .l2norm import l2norm_310p
__all__ = [
"fused_gdn_gating_pytorch",
"fused_recurrent_gated_delta_rule_pytorch",
"chunk_gated_delta_rule_310",
"l2norm_310p",
]

View File

@@ -0,0 +1,586 @@
#
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# mypy: ignore-errors
from __future__ import annotations
import torch
import torch.nn.functional as F
from vllm_ascend._310p.ops.fla.l2norm import l2norm_310p
CHUNK_SIZE = 64
def _expand_qk_to_v_heads(x: torch.Tensor, num_v_heads: int) -> torch.Tensor:
"""
Expand q/k heads to match v heads for grouped-value-attention semantics.
x: [L, Hqk, D] -> [L, Hv, D]
"""
h_qk = x.shape[-2]
if h_qk == num_v_heads:
return x
if num_v_heads % h_qk != 0:
raise ValueError(f"Invalid grouped heads: Hqk={h_qk}, Hv={num_v_heads}.")
group_size = num_v_heads // h_qk
return x.repeat_interleave(group_size, dim=-2)
def _iter_seq_ranges(batch_size: int, seq_len: int, cu_seqlens: torch.Tensor | None) -> list[tuple[int, int, int]]:
if cu_seqlens is None:
return [(i, 0, seq_len) for i in range(batch_size)]
return [(i, int(cu_seqlens[i].item()), int(cu_seqlens[i + 1].item())) for i in range(len(cu_seqlens) - 1)]
def _normalize_chunk_inputs(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
cu_seqlens: torch.Tensor | None,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, bool]:
"""
Normalize inputs to [B, T, H, D] / [B, T, H] while preserving TND support.
Returns normalized tensors and a flag indicating whether input was TND.
"""
input_was_tnd = False
if q.ndim == 3:
if cu_seqlens is None:
raise ValueError("TND inputs require `cu_seqlens` for variable-length layout.")
if k.ndim != 3 or v.ndim != 3:
raise ValueError("When q is TND, k and v must also be TND.")
if g.ndim != 2 or beta.ndim != 2:
raise ValueError("When q is TND, g and beta must be shape [T, H].")
q = q.unsqueeze(0)
k = k.unsqueeze(0)
v = v.unsqueeze(0)
g = g.unsqueeze(0)
beta = beta.unsqueeze(0)
input_was_tnd = True
elif q.ndim == 4:
if k.ndim != 4 or v.ndim != 4:
raise ValueError("When q is 4D, k and v must also be 4D.")
if g.ndim != 3 or beta.ndim != 3:
raise ValueError("When q is 4D, g and beta must be shape [B, T, H].")
else:
raise ValueError(f"Unsupported q ndim={q.ndim}; expected 3D(TND) or 4D(BTHD).")
return q, k, v, g, beta, input_was_tnd
def _torch_chunk_gated_delta_rule_chunked(
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
chunk_size: int = CHUNK_SIZE,
scale: float | None = None,
initial_state: torch.Tensor | None = None,
output_final_state: bool = False,
use_qk_l2norm_in_kernel: bool = False,
) -> tuple[torch.Tensor, torch.Tensor | None]:
"""
Chunked torch implementation aligned with the Qwen3-Next torch path:
transformers/models/qwen3_next/modular_qwen3_next.py::torch_chunk_gated_delta_rule
Shapes:
query/key: [B, T, H, K]
value: [B, T, H, V]
g/beta: [B, T, H]
initial_state: [B, H, V, K]
"""
initial_dtype = query.dtype
if use_qk_l2norm_in_kernel:
query = l2norm_310p(query)
key = l2norm_310p(key)
query, key, value, beta, g = [
x.transpose(1, 2).contiguous().to(torch.float32) for x in (query, key, value, beta, g)
]
batch_size, num_heads, sequence_length, k_head_dim = key.shape
v_head_dim = value.shape[-1]
pad_size = (chunk_size - sequence_length % chunk_size) % chunk_size
query = F.pad(query, (0, 0, 0, pad_size))
key = F.pad(key, (0, 0, 0, pad_size))
value = F.pad(value, (0, 0, 0, pad_size))
beta = F.pad(beta, (0, pad_size))
g = F.pad(g, (0, pad_size))
total_sequence_length = sequence_length + pad_size
scale = query.shape[-1] ** -0.5 if scale is None else scale
query = query * scale
v_beta = value * beta.unsqueeze(-1)
k_beta = key * beta.unsqueeze(-1)
# reshape to chunks
query, key, value, k_beta, v_beta = [
x.reshape(x.shape[0], x.shape[1], -1, chunk_size, x.shape[-1]) for x in (query, key, value, k_beta, v_beta)
]
g = g.reshape(g.shape[0], g.shape[1], -1, chunk_size)
mask_diag = torch.triu(torch.ones(chunk_size, chunk_size, dtype=torch.bool, device=query.device), diagonal=0)
# chunk decay
g = g.cumsum(dim=-1)
decay_mask = ((g.unsqueeze(-1) - g.unsqueeze(-2)).tril().exp().float()).tril()
attn = -((k_beta @ key.transpose(-1, -2)) * decay_mask).masked_fill(mask_diag, 0)
for i in range(1, chunk_size):
row = attn[..., i, :i].clone()
sub = attn[..., :i, :i].clone()
attn[..., i, :i] = row + (row.unsqueeze(-1) * sub).sum(-2)
attn = attn + torch.eye(chunk_size, dtype=attn.dtype, device=attn.device)
value = attn @ v_beta
k_cumdecay = attn @ (k_beta * g.exp().unsqueeze(-1))
last_recurrent_state = (
torch.zeros(batch_size, num_heads, v_head_dim, k_head_dim, device=value.device, dtype=value.dtype)
if initial_state is None
else initial_state.to(value)
)
core_attn_out = torch.zeros_like(value)
mask_upper = torch.triu(torch.ones(chunk_size, chunk_size, dtype=torch.bool, device=query.device), diagonal=1)
# for each chunk
for i in range(0, total_sequence_length // chunk_size):
q_i, k_i, v_i = query[:, :, i], key[:, :, i], value[:, :, i]
attn_inter_chunk = (q_i @ k_i.transpose(-1, -2) * decay_mask[:, :, i]).masked_fill_(mask_upper, 0)
v_prime = k_cumdecay[:, :, i] @ last_recurrent_state.transpose(-1, -2)
v_new = v_i - v_prime
inter_state = (q_i * g[:, :, i, :, None].exp()) @ last_recurrent_state.transpose(-1, -2)
core_attn_out[:, :, i] = inter_state + attn_inter_chunk @ v_new
last_recurrent_state = last_recurrent_state * g[:, :, i, -1, None, None].exp() + v_new.transpose(-1, -2) @ (
k_i * (g[:, :, i, -1, None] - g[:, :, i]).exp()[..., None]
)
if not output_final_state:
last_recurrent_state = None
core_attn_out = core_attn_out.reshape(core_attn_out.shape[0], core_attn_out.shape[1], -1, core_attn_out.shape[-1])
core_attn_out = core_attn_out[:, :, :sequence_length]
core_attn_out = core_attn_out.transpose(1, 2).contiguous().to(initial_dtype)
return core_attn_out, last_recurrent_state
def _ceil_div(value: int, divisor: int) -> int:
return (value + divisor - 1) // divisor
def _require_ascend_chunk_ops(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> None:
ascend_ops = getattr(torch.ops, "_C_ascend", None)
if q.device.type != "npu" or ascend_ops is None:
raise RuntimeError("310P chunk_gated_delta_rule requires NPU AscendC kernels.")
if not (hasattr(ascend_ops, "chunk_gated_delta_rule_fwd_h") and hasattr(ascend_ops, "chunk_fwd_o")):
raise RuntimeError("Missing AscendC chunk-gdr ops: chunk_gated_delta_rule_fwd_h/chunk_fwd_o.")
if q.dtype != torch.float16 or k.dtype != q.dtype or v.dtype != q.dtype:
raise TypeError(f"q/k/v must share float16 dtype on 310P, got {q.dtype}, {k.dtype}, {v.dtype}.")
if v.shape[-1] < 128 or v.shape[-1] % 128 != 0:
raise ValueError(f"v head dim must be >=128 and a multiple of 128, got {v.shape[-1]}.")
def _maybe_l2norm(x: torch.Tensor, enabled: bool) -> torch.Tensor:
if not enabled:
return x
return l2norm_310p(x)
def _pad_bthd_to_chunk(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
chunk_size: int,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, list[tuple[int, int, int]], None]:
batch_size, seq_len = q.shape[:2]
padded_len = _ceil_div(seq_len, chunk_size) * chunk_size
pad_len = padded_len - seq_len
seq_ranges = [(batch_idx, 0, seq_len) for batch_idx in range(batch_size)]
if pad_len == 0:
return q, k, v, g, beta, seq_ranges, None
q_pad = q.new_zeros((batch_size, pad_len, *q.shape[2:]))
k_pad = k.new_zeros((batch_size, pad_len, *k.shape[2:]))
v_pad = v.new_zeros((batch_size, pad_len, *v.shape[2:]))
g_pad = g.new_zeros((batch_size, pad_len, g.shape[-1]))
beta_pad = beta.new_zeros((batch_size, pad_len, beta.shape[-1]))
return (
torch.cat((q, q_pad), dim=1),
torch.cat((k, k_pad), dim=1),
torch.cat((v, v_pad), dim=1),
torch.cat((g, g_pad), dim=1),
torch.cat((beta, beta_pad), dim=1),
seq_ranges,
None,
)
def _pad_varlen_to_chunk(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
cu_seqlens: torch.Tensor,
chunk_size: int,
) -> tuple[
torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, list[tuple[int, int, int]], torch.Tensor
]:
q_parts: list[torch.Tensor] = []
k_parts: list[torch.Tensor] = []
v_parts: list[torch.Tensor] = []
g_parts: list[torch.Tensor] = []
beta_parts: list[torch.Tensor] = []
seq_ranges: list[tuple[int, int, int]] = []
padded_cu = [0]
out_cursor = 0
for seq_idx in range(cu_seqlens.numel() - 1):
start = int(cu_seqlens[seq_idx].item())
end = int(cu_seqlens[seq_idx + 1].item())
seq_len = end - start
padded_len = _ceil_div(seq_len, chunk_size) * chunk_size if seq_len > 0 else 0
pad_len = padded_len - seq_len
q_seq = q[:, start:end]
k_seq = k[:, start:end]
v_seq = v[:, start:end]
g_seq = g[:, start:end]
beta_seq = beta[:, start:end]
if pad_len > 0:
q_seq = torch.cat((q_seq, q.new_zeros((1, pad_len, *q.shape[2:]))), dim=1)
k_seq = torch.cat((k_seq, k.new_zeros((1, pad_len, *k.shape[2:]))), dim=1)
v_seq = torch.cat((v_seq, v.new_zeros((1, pad_len, *v.shape[2:]))), dim=1)
g_seq = torch.cat((g_seq, g.new_zeros((1, pad_len, g.shape[-1]))), dim=1)
beta_seq = torch.cat((beta_seq, beta.new_zeros((1, pad_len, beta.shape[-1]))), dim=1)
q_parts.append(q_seq)
k_parts.append(k_seq)
v_parts.append(v_seq)
g_parts.append(g_seq)
beta_parts.append(beta_seq)
seq_ranges.append((0, out_cursor, out_cursor + seq_len))
out_cursor += seq_len
padded_cu.append(padded_cu[-1] + padded_len)
if q_parts:
q_padded = torch.cat(q_parts, dim=1)
k_padded = torch.cat(k_parts, dim=1)
v_padded = torch.cat(v_parts, dim=1)
g_padded = torch.cat(g_parts, dim=1)
beta_padded = torch.cat(beta_parts, dim=1)
else:
q_padded = q[:, :0]
k_padded = k[:, :0]
v_padded = v[:, :0]
g_padded = g[:, :0]
beta_padded = beta[:, :0]
cu_padded = torch.tensor(padded_cu, dtype=torch.int64, device=cu_seqlens.device)
return q_padded, k_padded, v_padded, g_padded, beta_padded, seq_ranges, cu_padded
def _prepare_chunk_indices_list(cu_seqlens: torch.Tensor, chunk_size: int) -> list[int]:
chunk_indices: list[int] = []
compact_seq_idx = 0
for seq_idx in range(cu_seqlens.numel() - 1):
seq_len = int(cu_seqlens[seq_idx + 1].item() - cu_seqlens[seq_idx].item())
num_chunks = _ceil_div(seq_len, chunk_size) if seq_len > 0 else 0
if num_chunks == 0:
continue
for chunk_idx in range(num_chunks):
chunk_indices.extend((compact_seq_idx, chunk_idx))
compact_seq_idx += 1
return chunk_indices
def _compute_kernel_inputs_from_torch_wy(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
chunk_size: int,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
"""Compute the original torch WY prefix and return AscendC kernel layout."""
batch_size, padded_tokens, _, k_dim = k.shape
num_v_heads = v.shape[2]
value_dim = v.shape[-1]
num_chunks = padded_tokens // chunk_size
q_kernel = q.transpose(1, 2).contiguous()
k_kernel = k.transpose(1, 2).contiguous()
key = _expand_qk_to_v_heads(k, num_v_heads).transpose(1, 2).contiguous().to(torch.float32)
value = v.transpose(1, 2).contiguous().to(torch.float32)
g = g.transpose(1, 2).contiguous().to(torch.float32)
beta = beta.transpose(1, 2).contiguous().to(torch.float32)
key = key.reshape(batch_size, num_v_heads, num_chunks, chunk_size, k_dim)
value = value.reshape(batch_size, num_v_heads, num_chunks, chunk_size, value_dim)
g = g.reshape(batch_size, num_v_heads, num_chunks, chunk_size).cumsum(dim=-1)
beta = beta.reshape(batch_size, num_v_heads, num_chunks, chunk_size)
lower_decay = (g.unsqueeze(-1) - g.unsqueeze(-2)).tril().exp().float().tril()
k_beta = key * beta.unsqueeze(-1)
attn = -(k_beta @ key.transpose(-1, -2) * lower_decay)
mask_diag = torch.triu(
torch.ones(chunk_size, chunk_size, dtype=torch.bool, device=k.device),
diagonal=0,
)
attn = attn.masked_fill(mask_diag, 0)
for row_idx in range(1, chunk_size):
row = attn[..., row_idx, :row_idx].clone()
sub = attn[..., :row_idx, :row_idx].clone()
attn[..., row_idx, :row_idx] = row + (row.unsqueeze(-1) * sub).sum(-2)
attn = attn + torch.eye(chunk_size, dtype=attn.dtype, device=attn.device)
value = attn @ (value * beta.unsqueeze(-1))
k_cumdecay = attn @ (k_beta * g.exp().unsqueeze(-1))
u_kernel = value.reshape(batch_size, num_v_heads, padded_tokens, value_dim).to(torch.float16).contiguous()
w_kernel = k_cumdecay.reshape(batch_size, num_v_heads, padded_tokens, k_dim).to(torch.float16).contiguous()
g_kernel = g.reshape(batch_size, num_v_heads, padded_tokens).contiguous()
return q_kernel, k_kernel, w_kernel, u_kernel, g_kernel
def _unpad_chunk_output(
out: torch.Tensor,
seq_ranges: list[tuple[int, int, int]],
total_tokens: int,
input_was_tnd: bool,
is_varlen: bool,
) -> torch.Tensor:
if is_varlen:
unpadded = out.new_empty((1, total_tokens, *out.shape[2:]))
padded_cursor = 0
for _, start, end in seq_ranges:
seq_len = end - start
padded_len = _ceil_div(seq_len, CHUNK_SIZE) * CHUNK_SIZE if seq_len > 0 else 0
if seq_len > 0:
unpadded[:, start:end] = out[:, padded_cursor : padded_cursor + seq_len]
padded_cursor += padded_len
return unpadded[0] if input_was_tnd else unpadded
seq_len = total_tokens
return out[:, :seq_len]
def chunk_gated_delta_rule_pytorch(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
scale: float | None = None,
initial_state: torch.Tensor | None = None,
output_final_state: bool = False,
cu_seqlens: torch.Tensor | None = None,
head_first: bool = False,
use_qk_l2norm_in_kernel: bool = False,
) -> tuple[torch.Tensor, torch.Tensor | None]:
"""
Reference-only implementation with vLLM-compatible interface.
Internal math follows Transformers torch_chunk_gated_delta_rule flow.
The 310P production path must use the AscendC chunk-gdr kernels and does
not fall back to this implementation.
"""
if head_first:
raise DeprecationWarning("head_first=True is not supported in the reference implementation.")
q, k, v, g, beta, input_was_tnd = _normalize_chunk_inputs(q, k, v, g, beta, cu_seqlens)
if cu_seqlens is not None and q.shape[0] != 1:
raise ValueError("Variable-length mode expects batch size B=1.")
batch_size, total_tokens, h_qk, k_dim = q.shape
h_v = v.shape[2]
v_dim = v.shape[-1]
if k.shape != q.shape:
raise ValueError("q and k shapes must match.")
if g.shape != beta.shape or g.shape[:2] != (batch_size, total_tokens) or g.shape[2] != h_v:
raise ValueError("g/beta must have shape [B, T, Hv] matching v.")
seq_ranges = _iter_seq_ranges(batch_size, total_tokens, cu_seqlens)
num_states = batch_size if cu_seqlens is None else len(cu_seqlens) - 1
if initial_state is not None:
states = initial_state.to(torch.float32).clone()
else:
states = torch.zeros(num_states, h_v, v_dim, k_dim, dtype=torch.float32, device=q.device)
out = torch.zeros_like(v)
for seq_idx, start, end in seq_ranges:
seq_len = end - start
if seq_len <= 0:
continue
b_idx = 0 if (cu_seqlens is not None and batch_size == 1) else seq_idx
q_seq = _expand_qk_to_v_heads(q[b_idx, start:end], h_v).unsqueeze(0)
k_seq = _expand_qk_to_v_heads(k[b_idx, start:end], h_v).unsqueeze(0)
v_seq = v[b_idx, start:end].unsqueeze(0)
g_seq = g[b_idx, start:end].unsqueeze(0)
beta_seq = beta[b_idx, start:end].unsqueeze(0)
init_seq_state = states[seq_idx].unsqueeze(0)
out_seq, final_state = _torch_chunk_gated_delta_rule_chunked(
query=q_seq,
key=k_seq,
value=v_seq,
g=g_seq,
beta=beta_seq,
chunk_size=CHUNK_SIZE,
scale=scale,
initial_state=init_seq_state,
output_final_state=True,
use_qk_l2norm_in_kernel=use_qk_l2norm_in_kernel,
)
out[b_idx, start:end] = out_seq[0]
states[seq_idx] = final_state[0]
if input_was_tnd:
out = out[0]
if output_final_state:
return out, states
return out, None
def chunk_gated_delta_rule_310(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
scale: float | None = None,
initial_state: torch.Tensor | None = None,
output_final_state: bool = False,
cu_seqlens: torch.Tensor | None = None,
head_first: bool = False,
use_qk_l2norm_in_kernel: bool = False,
) -> tuple[torch.Tensor, torch.Tensor | None]:
"""310P chunk GDN path backed by AscendC fwd_h/fwd_o kernels.
Triton is unavailable on 310P, so the local WY preparation is done with
torch ops and the inter-chunk state/output matmuls are delegated to the
custom AscendC kernels.
"""
if head_first:
raise DeprecationWarning("head_first=True is not supported in 310P chunk path.")
if g is None or beta is None:
raise RuntimeError("g and beta are required for the AscendC chunk-gdr path.")
q, k, v, g, beta, input_was_tnd = _normalize_chunk_inputs(q, k, v, g, beta, cu_seqlens)
if cu_seqlens is not None and q.shape[0] != 1:
raise ValueError("Variable-length mode expects batch size B=1.")
if k.shape != q.shape:
raise ValueError("q and k shapes must match.")
if g.shape != beta.shape or g.shape[:2] != q.shape[:2] or g.shape[2] != v.shape[2]:
raise ValueError("g/beta must have shape [B, T, Hv] matching v.")
_require_ascend_chunk_ops(q, k, v)
q = _maybe_l2norm(q, use_qk_l2norm_in_kernel)
k = _maybe_l2norm(k, use_qk_l2norm_in_kernel)
original_tokens = v.shape[1]
if cu_seqlens is None:
q_pad, k_pad, v_pad, g_pad, beta_pad, seq_ranges, cu_kernel = _pad_bthd_to_chunk(q, k, v, g, beta, CHUNK_SIZE)
cu_list = None
chunk_indices_list = None
num_states = q.shape[0]
else:
q_pad, k_pad, v_pad, g_pad, beta_pad, seq_ranges, cu_kernel = _pad_varlen_to_chunk(
q, k, v, g, beta, cu_seqlens.to(torch.int64).cpu(), CHUNK_SIZE
)
assert cu_kernel is not None
cu_list = cu_kernel.tolist()
chunk_indices_list = _prepare_chunk_indices_list(cu_kernel, CHUNK_SIZE)
num_states = cu_seqlens.numel() - 1
expected_state_shape = (num_states, v.shape[2], v.shape[-1], k.shape[-1])
if initial_state is not None:
if initial_state.device != q.device:
raise RuntimeError(f"initial_state must be on {q.device}, got {initial_state.device}.")
if tuple(initial_state.shape) != expected_state_shape:
raise ValueError(f"initial_state must have shape {expected_state_shape}, got {tuple(initial_state.shape)}.")
if q_pad.shape[1] == 0:
empty_out = v.new_empty((0, *v.shape[2:])) if input_was_tnd else v.new_empty(v.shape)
final_state = initial_state if output_final_state else None
return empty_out, final_state
scale = k.shape[-1] ** -0.5 if scale is None else scale
q_kernel, k_kernel, w_kernel, u_kernel, g_kernel = _compute_kernel_inputs_from_torch_wy(
q_pad, k_pad, v_pad, g_pad, beta_pad, CHUNK_SIZE
)
if initial_state is None:
state = torch.zeros(
num_states,
v.shape[2],
v.shape[-1],
k.shape[-1],
dtype=torch.float32,
device=v.device,
)
else:
state = initial_state
state_kernel = state.transpose(-1, -2).contiguous()
h, v_new, final_state_kernel = torch.ops._C_ascend.chunk_gated_delta_rule_fwd_h(
k_kernel,
w_kernel,
u_kernel,
g=g_kernel,
gk=None,
initial_state=state_kernel,
output_final_state=output_final_state,
chunk_size=CHUNK_SIZE,
save_new_value=True,
cu_seqlens=cu_list,
chunk_indices=chunk_indices_list,
use_exp2=False,
transpose_state_layout=False,
)
o_kernel = torch.ops._C_ascend.chunk_fwd_o(
q_kernel,
k_kernel,
v_new,
h,
scale,
g=g_kernel,
g_gamma=None,
cu_seqlens=cu_list,
chunk_indices=chunk_indices_list,
chunk_size=CHUNK_SIZE,
transpose_state_layout=False,
)
out = o_kernel.transpose(1, 2).contiguous().to(v.dtype)
out = _unpad_chunk_output(out, seq_ranges, original_tokens, input_was_tnd, cu_seqlens is not None)
if not output_final_state:
return out, None
final_state = final_state_kernel.transpose(-1, -2).contiguous()
return out, final_state

View File

@@ -0,0 +1,57 @@
import numpy as np
import torch
from vllm.v1.attention.backends.utils import PAD_SLOT_ID
def compute_causal_conv1d_metadata(
query_start_loc_p_cpu: torch.Tensor,
*,
device: torch.device,
):
assert query_start_loc_p_cpu.device.type == "cpu"
seqlens = query_start_loc_p_cpu.diff()
nums_dict: dict[int, dict[str, object]] = {}
batch_ptr = None
token_chunk_offset_ptr = None
batch_ptr_cpu = None
token_chunk_offset_ptr_cpu = None
for BLOCK_M in [8]:
nums = -(-seqlens // BLOCK_M)
nums_dict[BLOCK_M] = {}
nums_dict[BLOCK_M]["nums"] = nums
nums_dict[BLOCK_M]["tot"] = nums.sum().item()
mlist = torch.from_numpy(np.repeat(np.arange(len(nums)), nums.numpy()))
nums_dict[BLOCK_M]["mlist"] = mlist
mlist_len = len(mlist)
nums_dict[BLOCK_M]["mlist_len"] = mlist_len
MAX_NUM_PROGRAMS = max(1024, mlist_len) * 2
offset_items: list[int] = []
for idx, num in enumerate(nums):
offset_items.extend(range(num))
offsetlist = torch.tensor(offset_items, dtype=torch.int32)
if batch_ptr is None or batch_ptr.numel() < MAX_NUM_PROGRAMS:
batch_ptr_cpu = torch.full((MAX_NUM_PROGRAMS,), PAD_SLOT_ID, dtype=torch.int32)
token_chunk_offset_ptr_cpu = torch.full((MAX_NUM_PROGRAMS,), PAD_SLOT_ID, dtype=torch.int32)
if device.type == "cpu":
batch_ptr = batch_ptr_cpu
token_chunk_offset_ptr = token_chunk_offset_ptr_cpu
else:
batch_ptr = batch_ptr_cpu.to(device, non_blocking=False)
token_chunk_offset_ptr = token_chunk_offset_ptr_cpu.to(device, non_blocking=False)
else:
batch_ptr_cpu.fill_(PAD_SLOT_ID)
token_chunk_offset_ptr_cpu.fill_(PAD_SLOT_ID)
batch_ptr_cpu[:mlist_len].copy_(mlist.to(torch.int32))
token_chunk_offset_ptr_cpu[:mlist_len].copy_(offsetlist)
if device.type != "cpu":
batch_ptr.copy_(batch_ptr_cpu, non_blocking=True)
token_chunk_offset_ptr.copy_(token_chunk_offset_ptr_cpu, non_blocking=True)
nums_dict[BLOCK_M]["batch_ptr"] = batch_ptr
nums_dict[BLOCK_M]["token_chunk_offset_ptr"] = token_chunk_offset_ptr
return nums_dict, batch_ptr, token_chunk_offset_ptr

View File

@@ -0,0 +1,62 @@
import torch
def fused_gdn_gating_pytorch(
A_log: torch.Tensor,
a: torch.Tensor,
b: torch.Tensor,
dt_bias: torch.Tensor,
beta: float = 1.0,
threshold: float = 20.0,
) -> tuple[torch.Tensor, torch.Tensor]:
"""
PyTorch implementation of fused_gdn_gating.
This is a fallback implementation for 310P without Triton support.
Args:
A_log: Log of A parameter, shape [num_heads]
a: a parameter, shape [batch, num_heads]
b: b parameter, shape [batch, num_heads]
dt_bias: dt bias, shape [num_heads]
beta: softplus beta parameter
threshold: softplus threshold parameter
Returns:
g: gating parameter, shape [1, batch, num_heads]
beta_output: sigmoid(b), shape [1, batch, num_heads]
"""
batch, num_heads = a.shape
del num_heads
# Keep nonlinear gating math in fp32 for stability.
compute_dtype = torch.float32
A_log_f = A_log.to(compute_dtype)
a_f = a.to(compute_dtype)
b_f = b.to(compute_dtype)
dt_bias_f = dt_bias.to(compute_dtype)
# Expand A_log and dt_bias to match a shape.
A_log_expanded = A_log_f.unsqueeze(0).expand(batch, -1)
dt_bias_expanded = dt_bias_f.unsqueeze(0).expand(batch, -1)
# Compute x = a + dt_bias.
x = a_f + dt_bias_expanded
# Compute softplus(x).
beta_x = beta * x
softplus_x = torch.where(
beta_x <= threshold,
(1.0 / beta) * torch.log1p(torch.exp(beta_x)),
x,
)
# Compute g = -exp(A_log) * softplus(x).
g = -torch.exp(A_log_expanded) * softplus_x
# Add sequence dimension.
g = g.unsqueeze(0)
# Match Triton kernel: sigmoid in fp32, then cast to input b dtype.
beta_output = torch.sigmoid(b_f).to(b.dtype)
beta_output = beta_output.unsqueeze(0)
return g, beta_output

View File

@@ -0,0 +1,227 @@
import torch
from vllm_ascend._310p.ops.fla.l2norm import l2norm_310p
def _maybe_l2norm(x: torch.Tensor, enabled: bool) -> torch.Tensor:
if not enabled:
return x
return l2norm_310p(x)
def _expand_to_hv(x: torch.Tensor, hv: int) -> torch.Tensor:
"""Expand [H, ...] to [HV, ...] for grouped-value-attention semantics."""
h = x.shape[0]
if h == hv:
return x
if hv % h != 0:
raise ValueError(f"Cannot expand head dim from {h} to {hv}.")
return x.repeat_interleave(hv // h, dim=0)
def _infer_num_states(
default_n: int,
initial_state: torch.Tensor | None,
ssm_state_indices: torch.Tensor | None,
) -> int:
if initial_state is not None:
return initial_state.shape[0]
if ssm_state_indices is None:
return default_n
nonneg = ssm_state_indices[ssm_state_indices >= 0]
if nonneg.numel() == 0:
return default_n
return int(nonneg.max().item()) + 1
def _state_index(
seq_idx: int,
tok_idx: int,
ssm_state_indices: torch.Tensor | None,
) -> int:
if ssm_state_indices is None:
return seq_idx
if ssm_state_indices.ndim == 1:
return int(ssm_state_indices[seq_idx].item())
return int(ssm_state_indices[seq_idx, tok_idx].item())
def _run_recurrent_gated_delta_rule(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor | None,
beta: torch.Tensor | None,
states: torch.Tensor,
scale: float,
cu_seqlens: torch.Tensor | None,
ssm_state_indices: torch.Tensor | None,
num_accepted_tokens: torch.Tensor | None,
use_initial_state: bool,
use_qk_l2norm_in_kernel: bool,
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Reference PyTorch recurrence for GDN delta rule.
Shapes follow fla.ops conventions:
q,k: [B, T, H, K]
v: [B, T, HV, V]
g,beta: [B, T, HV] (beta may also be [B, T, HV, V])
states: [N_state, HV, V, K]
"""
B, T, _, Kdim = k.shape
HV = v.shape[2]
Vdim = v.shape[-1]
if cu_seqlens is not None and B != 1:
raise ValueError("Variable-length mode expects batch size B=1.")
out = torch.zeros_like(v)
if cu_seqlens is None:
seq_ranges = [(i, 0, T) for i in range(B)]
else:
n_seq = len(cu_seqlens) - 1
seq_ranges = [
(
i,
int(cu_seqlens[i].item()),
int(cu_seqlens[i + 1].item()),
)
for i in range(n_seq)
]
for seq_idx, start, end in seq_ranges:
seq_len = end - start
if seq_len <= 0:
continue
accepted = None
if num_accepted_tokens is not None:
accepted = int(num_accepted_tokens[seq_idx].item())
seq_len = min(seq_len, accepted)
if seq_len <= 0:
continue
if use_initial_state:
if ssm_state_indices is None:
init_state_idx = seq_idx
else:
init_tok = (accepted - 1) if accepted is not None else 0
init_state_idx = _state_index(seq_idx, init_tok, ssm_state_indices)
if init_state_idx < 0:
# Match triton behavior for invalid PAD_SLOT_ID in continuous batching.
continue
if init_state_idx >= states.shape[0]:
raise IndexError(f"state_idx {init_state_idx} out of range for states size {states.shape[0]}")
h_t = states[init_state_idx].to(torch.float32)
else:
h_t = torch.zeros(HV, Vdim, Kdim, dtype=torch.float32, device=q.device)
for rel_t in range(seq_len):
tok = start + rel_t
if cu_seqlens is None:
q_t = q[seq_idx, tok]
k_t = k[seq_idx, tok]
v_t = v[seq_idx, tok]
g_t = g[seq_idx, tok] if g is not None else None
beta_t = beta[seq_idx, tok] if beta is not None else None
else:
q_t = q[0, tok]
k_t = k[0, tok]
v_t = v[0, tok]
g_t = g[0, tok] if g is not None else None
beta_t = beta[0, tok] if beta is not None else None
# Match Triton kernel math: load to fp32 first, then apply l2norm.
q_t = q_t.to(torch.float32)
k_t = k_t.to(torch.float32)
q_t = _maybe_l2norm(q_t, use_qk_l2norm_in_kernel)
k_t = _maybe_l2norm(k_t, use_qk_l2norm_in_kernel)
v_t = v_t.to(torch.float32)
q_t = q_t * scale
q_hv = _expand_to_hv(q_t, HV)
k_hv = _expand_to_hv(k_t, HV)
if g_t is not None:
g_t = g_t.to(torch.float32)
if g_t.ndim == 0:
g_t = g_t.expand(HV)
elif g_t.shape[0] != HV:
g_t = _expand_to_hv(g_t.unsqueeze(-1), HV).squeeze(-1)
h_t = h_t * torch.exp(g_t).view(HV, 1, 1)
v_t = v_t - torch.sum(h_t * k_hv.unsqueeze(-2), dim=-1)
if beta_t is not None:
beta_t = beta_t.to(torch.float32)
if beta_t.ndim == 1:
if beta_t.shape[0] != HV:
beta_t = _expand_to_hv(beta_t.unsqueeze(-1), HV).squeeze(-1)
v_t = v_t * beta_t.view(HV, 1)
else:
if beta_t.shape[0] != HV:
beta_t = _expand_to_hv(beta_t, HV)
v_t = v_t * beta_t
h_t = h_t + v_t.unsqueeze(-1) * k_hv.unsqueeze(-2)
o_t = torch.sum(h_t * q_hv.unsqueeze(-2), dim=-1)
if cu_seqlens is None:
out[seq_idx, tok] = o_t.to(out.dtype)
else:
out[0, tok] = o_t.to(out.dtype)
state_idx = _state_index(seq_idx, rel_t, ssm_state_indices)
if state_idx >= 0:
if state_idx >= states.shape[0]:
raise IndexError(f"state_idx {state_idx} out of range for states size {states.shape[0]}")
states[state_idx] = h_t.to(states.dtype)
return out, states
def fused_recurrent_gated_delta_rule_pytorch(
q,
k,
v,
g,
beta,
initial_state=None,
inplace_final_state=False,
cu_seqlens=None,
ssm_state_indices=None,
num_accepted_tokens=None,
use_qk_l2norm_in_kernel=False,
):
"""PyTorch fallback for fused_recurrent_gated_delta_rule."""
B, _, _, Kdim = k.shape
HV = v.shape[2]
Vdim = v.shape[-1]
N = B if cu_seqlens is None else len(cu_seqlens) - 1
n_states = _infer_num_states(N, initial_state, ssm_state_indices)
if initial_state is not None:
states = initial_state if inplace_final_state else initial_state.clone()
else:
states = torch.zeros(n_states, HV, Vdim, Kdim, dtype=q.dtype, device=q.device)
scale = Kdim**-0.5
out, states = _run_recurrent_gated_delta_rule(
q=q,
k=k,
v=v,
g=g,
beta=beta,
states=states,
scale=scale,
cu_seqlens=cu_seqlens,
ssm_state_indices=ssm_state_indices,
num_accepted_tokens=num_accepted_tokens,
use_initial_state=initial_state is not None,
use_qk_l2norm_in_kernel=use_qk_l2norm_in_kernel,
)
return out, states

View File

@@ -0,0 +1,429 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# from collections.abc import Iterable
# mypy: ignore-errors
import torch
from vllm.forward_context import get_forward_context
from vllm.model_executor.layers.mamba.gdn.base import GatedDeltaNetAttention
from vllm.v1.attention.backend import AttentionMetadata # type: ignore
from vllm.v1.attention.backends.gdn_attn import GDNAttentionMetadata
from vllm.v1.attention.backends.utils import PAD_SLOT_ID
from vllm_ascend._310p.ops.fla.chunk_gated_delta_rule import chunk_gated_delta_rule_310
from vllm_ascend._310p.ops.fla.fused_gdn_gating import fused_gdn_gating_pytorch
from vllm_ascend._310p.ops.fla.l2norm import l2norm_310p
from vllm_ascend.ascend_forward_context import _EXTRA_CTX
from vllm_ascend.attention.utils import maybe_save_kv_layer_to_connector
from vllm_ascend.utils import enable_sp
def _zero_padded_tokens(
tensor: torch.Tensor,
valid_tokens: torch.Tensor,
token_dim: int,
) -> torch.Tensor:
if tensor.numel() == 0:
return tensor
token_count = tensor.shape[token_dim]
if token_count == 0:
return tensor
positions = torch.arange(
token_count,
device=tensor.device,
dtype=valid_tokens.dtype,
)
valid_mask = positions < valid_tokens.to(device=tensor.device)
mask_shape = [1] * tensor.ndim
mask_shape[token_dim] = token_count
return tensor * valid_mask.reshape(mask_shape).to(dtype=tensor.dtype)
def _flatten_state_indices(
ssm_state_indices: torch.Tensor,
cu_seqlens: torch.Tensor,
total_tokens: int,
) -> torch.Tensor:
if ssm_state_indices.ndim == 1:
return ssm_state_indices[:total_tokens].to(torch.int32).contiguous()
num_seqs = (cu_seqlens[1:] - cu_seqlens[:-1]).shape[0]
seq_lens = cu_seqlens[1 : num_seqs + 1] - cu_seqlens[:num_seqs]
ssm_state_indices = ssm_state_indices[:num_seqs]
# Uniform spec-decode ACL graph uses fixed q_len per request; reshape avoids
# NPU masked_select which breaks stream capture (aclnnMaskedSelect / 107027).
if _EXTRA_CTX.capturing or (seq_lens.numel() > 0 and torch.all(seq_lens == seq_lens[0])):
q_per_seq = ssm_state_indices.shape[1]
flat = ssm_state_indices[:, :q_per_seq].reshape(-1)
return flat[:total_tokens].to(torch.int32).contiguous()
# Eager mixed batches with variable seq_lens: compact on CPU, copy back async.
ssm_cpu = ssm_state_indices.cpu()
seq_lens_cpu = seq_lens.cpu()
q_per_seq = ssm_cpu.shape[1]
positions = torch.arange(q_per_seq)
valid = positions.unsqueeze(0) < seq_lens_cpu.unsqueeze(1)
flat_cpu = ssm_cpu.masked_select(valid).to(torch.int32).contiguous()[:total_tokens]
if not flat_cpu.is_pinned:
flat_cpu = flat_cpu.pin_memory()
flat_dev = torch.empty(flat_cpu.numel(), dtype=torch.int32, device=ssm_state_indices.device)
flat_dev.copy_(flat_cpu, non_blocking=True)
return flat_dev.contiguous()
def _mask_padded_recurrent_accepted_tokens(
num_accepted_tokens: torch.Tensor,
actual_seq_lengths: torch.Tensor,
) -> torch.Tensor:
accepted_tokens = num_accepted_tokens[: actual_seq_lengths.shape[0]].to(torch.int32).contiguous()
return torch.where(
actual_seq_lengths > 0,
accepted_tokens,
torch.zeros_like(accepted_tokens),
).contiguous()
def npu_recurrent_gated_delta_rule_310(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor | None,
beta: torch.Tensor,
state: torch.Tensor,
cu_seqlens: torch.Tensor,
ssm_state_indices: torch.Tensor,
num_accepted_tokens: torch.Tensor | None = None,
use_qk_l2norm_in_kernel: bool = True,
) -> torch.Tensor:
if use_qk_l2norm_in_kernel:
q = l2norm_310p(q)
k = l2norm_310p(k)
total_tokens = v.shape[1]
flat_state_indices = _flatten_state_indices(ssm_state_indices, cu_seqlens, total_tokens)
actual_seq_lengths = (cu_seqlens[1:] - cu_seqlens[:-1]).to(torch.int32).contiguous()
flat_state_indices = torch.clamp_min(
flat_state_indices,
0,
).contiguous()
accepted_tokens = None
if num_accepted_tokens is not None:
accepted_tokens = _mask_padded_recurrent_accepted_tokens(
num_accepted_tokens,
actual_seq_lengths,
)
out = torch.ops._C_ascend.npu_recurrent_gated_delta_rule_310(
query=q.squeeze(0).to(torch.float16).contiguous(),
key=k.squeeze(0).to(torch.float16).contiguous(),
value=v.squeeze(0).to(torch.float16).contiguous(),
g=None if g is None else g.squeeze(0).to(torch.float32).contiguous(),
gk=None,
beta=beta.squeeze(0).to(torch.float16).contiguous(),
state=state,
actual_seq_lengths=actual_seq_lengths,
ssm_state_indices=flat_state_indices,
num_accepted_tokens=accepted_tokens,
scale_value=k.shape[-1] ** -0.5,
).unsqueeze(0)
return out
def _310p_get_state_dtype(self) -> tuple[torch.dtype, torch.dtype]:
conv_state_dtype, _ = _original_get_state_dtype(self)
return conv_state_dtype, torch.float16
_original_get_state_dtype = GatedDeltaNetAttention.get_state_dtype
def _merge_spec_and_non_spec_outputs_310(
core_attn_out: torch.Tensor,
num_actual_tokens: int,
spec_token_indx: torch.Tensor,
non_spec_token_indx: torch.Tensor,
core_attn_out_spec: torch.Tensor,
core_attn_out_non_spec: torch.Tensor,
) -> None:
"""Merge spec/non-spec GDN outputs back into the batch layout.
Avoid NPU ``index_copy_`` (IndexPutV2) which fails on some layouts; use
direct indexing instead. Validate lengths so mixed prefill+spec batches
do not pass mismatched tensors from spec ops.
"""
spec_out = core_attn_out_spec.squeeze(0)
non_spec_out = core_attn_out_non_spec.squeeze(0)
n_spec = spec_token_indx.numel()
n_non_spec = non_spec_token_indx.numel()
if spec_out.shape[0] != n_spec:
raise RuntimeError(f"GDN spec output length {spec_out.shape[0]} != spec_token_indx {n_spec}")
if non_spec_out.shape[0] != n_non_spec:
raise RuntimeError(f"GDN non-spec output length {non_spec_out.shape[0]} != non_spec_token_indx {n_non_spec}")
out = core_attn_out[:num_actual_tokens]
out[spec_token_indx] = spec_out
out[non_spec_token_indx] = non_spec_out
class AscendGatedDeltaNetAttention310(GatedDeltaNetAttention):
get_state_dtype = _310p_get_state_dtype
def get_attn_backend(self):
from vllm_ascend._310p.ops.gdn_attn_builder_310 import (
AscendGDNAttentionBackend310,
)
return AscendGDNAttentionBackend310
def _forward_core(
self,
mixed_qkv: torch.Tensor,
b: torch.Tensor,
a: torch.Tensor,
core_attn_out: torch.Tensor,
):
# Core attention computation (called by custom op).
# NOTE: The processing logic of Qwen3_5GatedDeltaNet is the same as Qwen3NextGatedDeltaNet.
# However, because the ops `torch_npu.npu_recurrent_gated_delta_rule`
# currently does not support `ssm_state` inputs in float32 format,
# we temporarily retain the current _forward_core implementation.
# Once the ops supports float32 `ssm_state`, this patch should be removed.
forward_context = get_forward_context()
attn_metadata: AttentionMetadata = forward_context.attn_metadata
if attn_metadata is None:
# V1 profile run
return
assert isinstance(attn_metadata, dict)
attn_metadata = attn_metadata[self.prefix]
assert isinstance(attn_metadata, GDNAttentionMetadata)
has_initial_state = attn_metadata.has_initial_state
spec_query_start_loc = attn_metadata.spec_query_start_loc
non_spec_query_start_loc = attn_metadata.non_spec_query_start_loc
spec_sequence_masks = attn_metadata.spec_sequence_masks
spec_token_indx = attn_metadata.spec_token_indx
non_spec_token_indx = attn_metadata.non_spec_token_indx
spec_state_indices_tensor = attn_metadata.spec_state_indices_tensor # noqa: E501
non_spec_state_indices_tensor = attn_metadata.non_spec_state_indices_tensor # noqa: E501
self_kv_cache = self.kv_cache
conv_state = self_kv_cache[0]
ssm_state = self_kv_cache[1]
num_actual_tokens = attn_metadata.num_actual_tokens
if not enable_sp():
mixed_qkv = mixed_qkv[:num_actual_tokens]
b = b[:num_actual_tokens]
a = a[:num_actual_tokens]
# 1. Convolution sequence transformation
conv_weights = self.conv1d.weight.view(self.conv1d.weight.size(0), self.conv1d.weight.size(2)).transpose(0, 1)
if spec_sequence_masks is not None:
if attn_metadata.num_prefills == 0 and attn_metadata.num_decodes == 0:
mixed_qkv_spec = mixed_qkv
mixed_qkv_non_spec = None
else:
mixed_qkv_spec = mixed_qkv.index_select(0, spec_token_indx)
mixed_qkv_non_spec = mixed_qkv.index_select(0, non_spec_token_indx)
else:
mixed_qkv_spec = None
mixed_qkv_non_spec = mixed_qkv
activation_num = 1 if self.activation else 0
# 1.1: Process the multi-query part
if spec_sequence_masks is not None:
spec_causal_conv1d_meta = attn_metadata.spec_decode_metadata.spec_causal_conv1d
spec_query_start_loc_device = spec_causal_conv1d_meta.query_start_loc
uniform_spec_only = attn_metadata.num_prefills == 0 and attn_metadata.num_decodes == 0
# The final entry remains the runtime token count even when
# graph metadata includes padded requests.
spec_valid_tokens = spec_query_start_loc_device[-1]
if uniform_spec_only:
mixed_qkv_spec = _zero_padded_tokens(
mixed_qkv_spec,
spec_valid_tokens,
token_dim=0,
)
mixed_qkv_spec = torch.ops._C_ascend.npu_causal_conv1d_310(
mixed_qkv_spec,
conv_weights,
bias=self.conv1d.bias,
conv_states=conv_state,
query_start_loc=spec_query_start_loc_device,
cache_indices=spec_causal_conv1d_meta.cache_indices,
initial_state_mode=None,
num_accepted_tokens=spec_causal_conv1d_meta.num_accepted_tokens,
activation_mode=activation_num,
pad_slot_id=PAD_SLOT_ID,
run_mode=1,
)
# 1.2: Process the remaining part
if attn_metadata.num_prefills > 0:
if mixed_qkv_non_spec is not None:
mixed_qkv_non_spec = torch.ops._C_ascend.npu_causal_conv1d_310(
mixed_qkv_non_spec,
conv_weights,
bias=self.conv1d.bias,
conv_states=conv_state,
query_start_loc=non_spec_query_start_loc,
cache_indices=non_spec_state_indices_tensor,
initial_state_mode=has_initial_state,
num_accepted_tokens=None,
activation_mode=activation_num,
pad_slot_id=PAD_SLOT_ID,
run_mode=0,
)
elif attn_metadata.num_decodes > 0:
mixed_qkv_non_spec = torch.ops._C_ascend.npu_causal_conv1d_310(
mixed_qkv_non_spec,
conv_weights,
bias=self.conv1d.bias,
conv_states=conv_state,
query_start_loc=None,
cache_indices=non_spec_state_indices_tensor[: attn_metadata.num_actual_tokens],
initial_state_mode=None,
num_accepted_tokens=None,
activation_mode=activation_num,
pad_slot_id=PAD_SLOT_ID,
run_mode=1,
)
else:
mixed_qkv_non_spec = None
query_spec, key_spec, value_spec = self.rearrange_mixed_qkv(mixed_qkv_spec)
query_non_spec, key_non_spec, value_non_spec = self.rearrange_mixed_qkv(mixed_qkv_non_spec)
g, beta = fused_gdn_gating_pytorch(self.A_log, a, b, self.dt_bias)
if attn_metadata.num_prefills > 0 or spec_sequence_masks is not None:
if spec_sequence_masks is not None:
if attn_metadata.num_prefills == 0 and attn_metadata.num_decodes == 0:
g_spec = g
beta_spec = beta
g_non_spec = None
beta_non_spec = None
else:
g_spec = g.index_select(1, spec_token_indx)
beta_spec = beta.index_select(1, spec_token_indx)
g_non_spec = g.index_select(1, non_spec_token_indx)
beta_non_spec = beta.index_select(1, non_spec_token_indx)
else:
g_spec = None
beta_spec = None
g_non_spec = g
beta_non_spec = beta
# 2. Recurrent attention
# 2.1: Process the multi-query part
if spec_sequence_masks is not None:
core_attn_out_spec = npu_recurrent_gated_delta_rule_310(
q=query_spec,
k=key_spec,
v=value_spec,
g=g_spec,
beta=beta_spec,
state=ssm_state,
cu_seqlens=spec_query_start_loc[: attn_metadata.num_spec_decodes + 1],
ssm_state_indices=spec_state_indices_tensor,
num_accepted_tokens=spec_causal_conv1d_meta.num_accepted_tokens,
use_qk_l2norm_in_kernel=True,
)
else:
core_attn_out_spec = None
# 2.2: Process the remaining part
if attn_metadata.num_prefills > 0:
initial_state = ssm_state[non_spec_state_indices_tensor].contiguous()
initial_state[~has_initial_state, ...] = 0
(
core_attn_out_non_spec,
last_recurrent_state,
) = chunk_gated_delta_rule_310(
q=query_non_spec,
k=key_non_spec,
v=value_non_spec,
g=g_non_spec,
beta=beta_non_spec,
initial_state=initial_state,
output_final_state=True,
cu_seqlens=non_spec_query_start_loc,
head_first=False,
use_qk_l2norm_in_kernel=True,
)
# Init cache
ssm_state[non_spec_state_indices_tensor] = last_recurrent_state.to(ssm_state.dtype)
elif attn_metadata.num_decodes > 0:
core_attn_out_non_spec = npu_recurrent_gated_delta_rule_310(
q=query_non_spec,
k=key_non_spec,
v=value_non_spec,
g=g_non_spec,
beta=beta_non_spec,
state=ssm_state,
cu_seqlens=non_spec_query_start_loc[: attn_metadata.num_decodes + 1],
ssm_state_indices=non_spec_state_indices_tensor,
use_qk_l2norm_in_kernel=True,
)
else:
core_attn_out_non_spec = None
elif attn_metadata.num_decodes > 0:
core_attn_out_non_spec = npu_recurrent_gated_delta_rule_310(
q=query_non_spec,
k=key_non_spec,
v=value_non_spec,
g=g,
beta=beta,
state=ssm_state,
cu_seqlens=non_spec_query_start_loc,
ssm_state_indices=non_spec_state_indices_tensor,
use_qk_l2norm_in_kernel=True,
)
# 3. Merge core attention output
if spec_sequence_masks is not None and core_attn_out_non_spec is not None:
_merge_spec_and_non_spec_outputs_310(
core_attn_out,
num_actual_tokens,
spec_token_indx,
non_spec_token_indx,
core_attn_out_spec,
core_attn_out_non_spec,
)
elif spec_sequence_masks is not None:
if not enable_sp():
core_attn_out[:num_actual_tokens] = core_attn_out_spec.squeeze(0)
else:
core_attn_out[:num_actual_tokens] = core_attn_out_spec.squeeze(0)[:num_actual_tokens]
else:
if not enable_sp():
core_attn_out[:num_actual_tokens] = core_attn_out_non_spec.squeeze(0)
else:
core_attn_out[:num_actual_tokens] = core_attn_out_non_spec.squeeze(0)[:num_actual_tokens]
if spec_sequence_masks is not None and uniform_spec_only:
core_attn_out.copy_(
_zero_padded_tokens(
core_attn_out,
spec_valid_tokens,
token_dim=0,
)
)
maybe_save_kv_layer_to_connector("", [])

View File

@@ -0,0 +1,42 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# from collections.abc import Iterable
# mypy: ignore-errors
import torch
from vllm.model_executor.layers.fla.ops.index import prepare_lens
from vllm.model_executor.layers.fla.ops.utils import tensor_cache
@tensor_cache
def prepare_chunk_indices_310(cu_seqlens: torch.Tensor, chunk_size: int) -> torch.Tensor:
seq_lens = prepare_lens(cu_seqlens)
num_chunks = (seq_lens + chunk_size - 1) // chunk_size
indices_list = []
for n in num_chunks.tolist():
indices_list.append(torch.arange(n, device=cu_seqlens.device))
indices = torch.cat(indices_list)
return torch.stack([indices.eq(0).cumsum(0) - 1, indices], dim=1).to(cu_seqlens)
@tensor_cache
def prepare_chunk_offsets_310(cu_seqlens: torch.Tensor, chunk_size: int) -> torch.Tensor:
seq_lens = prepare_lens(cu_seqlens)
num_chunks = (seq_lens + chunk_size - 1) // chunk_size
return torch.cat([torch.tensor([0], device=cu_seqlens.device, dtype=cu_seqlens.dtype), num_chunks]).cumsum(dim=-1)

View File

@@ -0,0 +1,43 @@
#
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# mypy: ignore-errors
from __future__ import annotations
import math
import torch
import torch_npu
from vllm.model_executor.layers.fla.ops.utils import tensor_cache
@tensor_cache
def _l2norm_unit_weight(dim: int, dtype: torch.dtype, device: torch.device) -> torch.Tensor:
# RMSNorm with weight 1/sqrt(dim) matches L2 norm: x / sqrt(sum(x^2)).
return torch.full((dim,), 1.0 / math.sqrt(dim), dtype=dtype, device=device)
def l2norm_310p(x: torch.Tensor, eps: float = 1e-6) -> torch.Tensor:
"""L2-normalize the last dimension using the 310P NPU RMSNorm kernel."""
orig_shape = x.shape
dim = x.shape[-1]
x_2d = x.reshape(-1, dim).contiguous()
weight = _l2norm_unit_weight(dim, x.dtype, x.device)
# RMSNorm: y = x / sqrt(mean(x^2) + eps_rms) * weight
# With weight=1/sqrt(dim), this equals x / sqrt(sum(x^2) + dim * eps_rms).
# L2 norm needs y = x / sqrt(sum(x^2) + eps), so eps_rms = eps / dim.
y, _ = torch_npu.npu_rms_norm(x_2d, weight, eps / dim)
return y.reshape(orig_shape)