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

View File

@@ -0,0 +1,32 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
#
# 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.
# This file is a part of the vllm-ascend project.
#
import torch
import torch.nn.functional as F
import torch_npu
from vllm_ascend.ops.activation import AscendSiluAndMul
class AscendSiluAndMul310(AscendSiluAndMul):
def forward(self, x: torch.Tensor) -> torch.Tensor:
if x.shape[-1] % 32 == 0:
out = torch_npu.npu_swiglu(x)
else:
h = x.shape[-1] // 2
out = F.silu(x[..., :h]) * x[..., h:]
return out

View File

@@ -0,0 +1,297 @@
import torch
import torch.nn.functional as F
from vllm.v1.attention.backends.utils import PAD_SLOT_ID
def causal_conv1d_ref(
x: torch.Tensor,
weight: torch.Tensor,
bias: torch.Tensor | None = None,
initial_states: torch.Tensor | None = None,
return_final_states: bool = False,
final_states_out: torch.Tensor | None = None,
activation: str | None = "silu",
):
"""
PyTorch reference implementation of causal_conv1d.
Args:
x: (batch, dim, seqlen)
weight: (dim, width)
bias: (dim,)
initial_states: (batch, dim, width - 1)
final_states_out: (batch, dim, width - 1)
return_final_states: bool
activation: str
Returns:
out: (batch, dim, seqlen)
final_states_out: (batch, dim, width - 1) if return_final_states
"""
if activation not in [None, "silu", "swish"]:
raise NotImplementedError("activation must be None, silu, or swish")
dtype_in = x.dtype
x = x.to(weight.dtype)
seqlen = x.shape[-1]
dim, width = weight.shape
if initial_states is None:
out = F.conv1d(x, weight.unsqueeze(1), bias, padding=width - 1, groups=dim)
else:
x = torch.cat([initial_states, x], dim=-1)
out = F.conv1d(x, weight.unsqueeze(1), bias, padding=0, groups=dim)
out = out[..., :seqlen]
if return_final_states:
final_states = F.pad(x, (width - 1 - x.shape[-1], 0)).to(dtype_in)
if final_states_out is not None:
final_states_out.copy_(final_states)
else:
final_states_out = final_states
out = (out if activation is None else F.silu(out)).to(dtype=dtype_in)
return (out, None) if not return_final_states else (out, final_states_out)
def causal_conv1d_fn(
x: torch.Tensor,
weight: torch.Tensor,
bias: torch.Tensor | None = None,
activation: str | None = "silu",
conv_states: torch.Tensor | None = None,
has_initial_state: torch.Tensor | None = None,
cache_indices: torch.Tensor | None = None,
query_start_loc: torch.Tensor | None = None,
pad_slot_id: int = PAD_SLOT_ID,
):
"""
PyTorch implementation of causal_conv1d_fn for 310P.
Args:
x: (dim, cu_seq_len) for varlen
weight: (dim, width)
bias: (dim,)
activation: str
conv_states: (..., dim, width - 1)
has_initial_state: (batch) bool
cache_indices: (batch) int32
query_start_loc: (batch + 1) int32
pad_slot_id: int
Returns:
out: (batch, dim, seqlen)
"""
if activation not in [None, "silu", "swish"]:
raise NotImplementedError("activation must be None, silu, or swish")
if query_start_loc is None:
raise RuntimeError("causal_conv1d_fn requires query_start_loc for varlen inputs.")
if cache_indices is None:
raise RuntimeError("causal_conv1d_fn requires cache_indices.")
if has_initial_state is None:
raise RuntimeError("causal_conv1d_fn requires has_initial_state.")
if conv_states is None:
raise RuntimeError("causal_conv1d_fn requires conv_states.")
if x.stride(-1) != 1:
x = x.contiguous()
bias = bias.contiguous() if bias is not None else None
# Normalize x to [dim, total_tokens]
if x.dim() == 3:
if x.shape[0] == 1:
x = x.squeeze(0)
elif x.shape[1] == 1:
x = x.squeeze(1).transpose(0, 1)
else:
raise RuntimeError(f"Unsupported x shape for causal_conv1d_fn: {tuple(x.shape)}")
if x.dim() != 2:
raise RuntimeError(f"Unsupported x ndim for causal_conv1d_fn: {x.dim()}")
feature_dim = x.shape[0]
if weight.shape[0] != feature_dim and weight.shape[1] == feature_dim:
weight = weight.transpose(0, 1)
weight = weight.contiguous()
dim, width = weight.shape
if dim != feature_dim:
raise RuntimeError(
f"causal_conv1d_fn: weight dim mismatch, x dim={feature_dim}, weight.shape={tuple(weight.shape)}"
)
state_len = width - 1
if conv_states.shape[-2] != dim and conv_states.shape[-1] == dim:
conv_states = conv_states.transpose(-1, -2)
if conv_states.shape[-2] != dim:
raise RuntimeError(
f"causal_conv1d_fn: conv_states dim mismatch, "
f"expected dim={dim}, conv_states.shape={tuple(conv_states.shape)}"
)
if conv_states.shape[-1] < state_len:
raise RuntimeError(f"causal_conv1d_fn: conv_states too short, need >= {state_len}, got {conv_states.shape[-1]}")
seqlens = (query_start_loc[1:] - query_start_loc[:-1]).tolist()
splits = torch.split(x, seqlens, dim=-1)
out_chunks = []
for i, x_s in enumerate(splits):
cache_idx = int(cache_indices[i].item())
if cache_idx == pad_slot_id:
continue
state = conv_states[cache_idx]
init_state = state[..., :state_len].unsqueeze(0) if bool(has_initial_state[i].item()) else None
out_ref, final_state = causal_conv1d_ref(
x_s.unsqueeze(0),
weight,
bias,
activation=activation,
return_final_states=True,
initial_states=init_state,
)
state[..., :state_len].copy_(final_state.squeeze(0))
out_chunks.append(out_ref.squeeze(0))
if not out_chunks:
return x.new_zeros((dim, 0))
return torch.cat(out_chunks, dim=-1)
def causal_conv1d_update(
x: torch.Tensor,
conv_state: torch.Tensor,
weight: torch.Tensor,
bias: torch.Tensor | None = None,
activation: bool | str | None = None,
conv_state_indices: torch.Tensor | None = None,
num_accepted_tokens: torch.Tensor | None = None,
query_start_loc: torch.Tensor | None = None,
pad_slot_id: int = PAD_SLOT_ID,
):
"""
PyTorch implementation of causal_conv1d_update for 310P.
Args:
x: Input tensor
conv_state: (..., dim, state_len)
weight: (dim, width)
bias: (dim,)
activation: str
conv_state_indices: (batch,) int32
num_accepted_tokens: (batch,) int32
query_start_loc: (batch + 1,) int32
pad_slot_id: int
Returns:
out: same shape as x
"""
if isinstance(activation, bool):
activation = "silu" if activation is True else None
elif activation is not None:
assert activation in ["silu", "swish"]
original_x_dtype = x.dtype
x = x.to(conv_state.dtype)
feature_dim = x.shape[-1] if query_start_loc is None else x.shape[1]
if weight.shape[0] != feature_dim and weight.shape[1] == feature_dim:
weight = weight.transpose(0, 1)
weight = weight.contiguous()
dim, width = weight.shape
if dim != feature_dim:
raise RuntimeError(
f"causal_conv1d_update: weight dim mismatch, feature_dim={feature_dim}, weight.shape={tuple(weight.shape)}"
)
if conv_state.shape[-2] != dim and conv_state.shape[-1] == dim:
# Accept both (..., dim, state_len) and (..., state_len, dim) inputs.
conv_state = conv_state.transpose(-1, -2)
if conv_state.shape[-2] != dim:
raise RuntimeError(
f"causal_conv1d_update: conv_state dim mismatch, "
f"expected dim={dim}, conv_state.shape={tuple(conv_state.shape)}"
)
state_len = width - 1
if conv_state.shape[-1] < state_len:
raise RuntimeError(
f"causal_conv1d_update: conv_state too short, need >= {state_len}, got {conv_state.shape[-1]}"
)
out = x.clone()
def _select_state(i: int) -> torch.Tensor | None:
if conv_state_indices is not None:
idx = int(conv_state_indices[i].item())
if idx == pad_slot_id:
return None
state = conv_state[idx]
else:
state = conv_state[i]
return state
def _run_one(seq_tokens: torch.Tensor, state: torch.Tensor) -> torch.Tensor:
# seq_tokens: [L, dim] -> [1, dim, L]
x_ref = seq_tokens.transpose(0, 1).unsqueeze(0)
init_state = state[..., :state_len].unsqueeze(0)
out_ref, final_state = causal_conv1d_ref(
x_ref,
weight,
bias,
initial_states=init_state,
return_final_states=True,
activation=activation,
)
state[..., :state_len].copy_(final_state.squeeze(0))
# [1, dim, L] -> [L, dim]
return out_ref.squeeze(0).transpose(0, 1)
if query_start_loc is None:
if x.dim() == 2:
batch = x.shape[0]
for i in range(batch):
state = _select_state(i)
if state is None:
continue
seq_tokens = x[i : i + 1]
if num_accepted_tokens is not None:
accepted = int(num_accepted_tokens[i].item())
if accepted <= 0:
continue
seq_tokens = seq_tokens[:accepted]
out_i = _run_one(seq_tokens, state)
out[i : i + out_i.shape[0]] = out_i
else:
batch = x.shape[0]
for i in range(batch):
state = _select_state(i)
if state is None:
continue
seq_tokens = x[i]
if num_accepted_tokens is not None:
accepted = int(num_accepted_tokens[i].item())
if accepted <= 0:
continue
seq_tokens = seq_tokens[:accepted]
out_i = _run_one(seq_tokens, state)
out[i, : out_i.shape[0]] = out_i
else:
assert conv_state_indices is not None
batch = conv_state_indices.size(0)
for i in range(batch):
start = int(query_start_loc[i].item())
end = int(query_start_loc[i + 1].item())
if end <= start:
continue
state = _select_state(i)
if state is None:
continue
seq_tokens = x[start:end]
if num_accepted_tokens is not None:
accepted = int(num_accepted_tokens[i].item())
if accepted <= 0:
continue
seq_tokens = seq_tokens[:accepted]
out_i = _run_one(seq_tokens, state)
out[start : start + out_i.shape[0]] = out_i
return out.to(original_x_dtype)

View File

@@ -0,0 +1,27 @@
#
# 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.
#
import torch
from vllm_ascend.ops.conv import AscendConv3dLayer
class AscendConv3dLayer310(AscendConv3dLayer):
def forward_oot(self, x: torch.Tensor) -> torch.Tensor:
# 310P should avoid the aclnn BatchMatMulV2 Conv3D path used by
# AscendConv3dLayer and keep vLLM's native Conv3d dispatch behavior.
return super().forward_native(x)

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)

View File

@@ -0,0 +1,249 @@
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
#
# 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.
"""310P RC GDN metadata builder.
This 310P-specific builder keeps the upstream RC-safe prefill metadata path
and adds ACL graph replay padding for decode / speculative decode metadata.
"""
from __future__ import annotations
import torch
from vllm.v1.attention.backend import CommonAttentionMetadata
from vllm.v1.attention.backends.gdn_attn import GDNAttentionMetadata
from vllm.v1.attention.backends.utils import NULL_BLOCK_ID
from vllm_ascend._310p.ops.fla.cumpute_causal_conv1d_metadata_310 import (
compute_causal_conv1d_metadata,
)
from vllm_ascend.ops.gdn_attn_builder import (
AscendGDNAttentionBackend,
AscendGDNAttentionMetadataBuilder,
)
class GDNAttentionMetadataBuilder310(AscendGDNAttentionMetadataBuilder):
"""310P overrides on top of :class:`AscendGDNAttentionMetadataBuilder`.
310P does not support Triton, so fallback metadata attachment is skipped.
For ACL graph replay, decode metadata is padded into fixed graph buffers.
"""
use_full_cuda_graph: bool
def _build_prefill_has_initial_state_and_causal_conv1d_meta(
self,
*,
common_attn_metadata: CommonAttentionMetadata,
context_lens_tensor: torch.Tensor,
num_prefills: int,
spec_sequence_masks_cpu: torch.Tensor | None,
non_spec_sequence_indices: torch.Tensor | None,
non_spec_query_start_loc_cpu: torch.Tensor | None,
query_start_loc: torch.Tensor,
) -> tuple[
torch.Tensor | None,
dict[int, dict[str, object]] | None,
torch.Tensor | None,
torch.Tensor | None,
]:
del common_attn_metadata, num_prefills
assert non_spec_query_start_loc_cpu is not None
has_initial_state = context_lens_tensor > 0
if spec_sequence_masks_cpu is not None:
assert non_spec_sequence_indices is not None
has_initial_state = torch.index_select(
has_initial_state,
0,
non_spec_sequence_indices,
)
nums_dict, batch_ptr, token_chunk_offset_ptr = compute_causal_conv1d_metadata(
non_spec_query_start_loc_cpu,
device=query_start_loc.device,
)
return (
has_initial_state,
nums_dict,
batch_ptr,
token_chunk_offset_ptr,
)
def _attach_non_spec_prefill_fallback_meta(
self,
attn_metadata: GDNAttentionMetadata,
common_attn_metadata: CommonAttentionMetadata,
non_spec_query_start_loc_cpu: torch.Tensor | None,
) -> GDNAttentionMetadata:
del common_attn_metadata, non_spec_query_start_loc_cpu
return attn_metadata
def _attach_spec_decode_fallback_meta(
self,
attn_metadata: GDNAttentionMetadata,
common_attn_metadata: CommonAttentionMetadata,
num_decode_draft_tokens_cpu: torch.Tensor | None,
) -> GDNAttentionMetadata:
del common_attn_metadata, num_decode_draft_tokens_cpu
return attn_metadata
def _attach_non_spec_decode_fallback_meta(
self,
attn_metadata: GDNAttentionMetadata,
common_attn_metadata: CommonAttentionMetadata,
num_decode_draft_tokens_cpu: torch.Tensor | None,
) -> GDNAttentionMetadata:
del common_attn_metadata, num_decode_draft_tokens_cpu
return attn_metadata
def _pad_spec_decode_metadata(
self,
attn_metadata: GDNAttentionMetadata,
graph_batch_size: int,
) -> None:
num_spec_decodes = attn_metadata.num_spec_decodes
spec_state_indices = attn_metadata.spec_state_indices_tensor
spec_sequence_masks = attn_metadata.spec_sequence_masks
spec_query_start_loc = attn_metadata.spec_query_start_loc
num_accepted_tokens = attn_metadata.num_accepted_tokens
assert spec_state_indices is not None
assert spec_sequence_masks is not None
assert spec_query_start_loc is not None
assert num_accepted_tokens is not None
self.spec_state_indices_tensor[:num_spec_decodes].copy_(
spec_state_indices,
non_blocking=True,
)
attn_metadata.spec_state_indices_tensor = self.spec_state_indices_tensor[:graph_batch_size]
attn_metadata.spec_state_indices_tensor[num_spec_decodes:].fill_(NULL_BLOCK_ID)
self.spec_sequence_masks[:num_spec_decodes].copy_(
spec_sequence_masks[:num_spec_decodes],
non_blocking=True,
)
attn_metadata.spec_sequence_masks = self.spec_sequence_masks[:graph_batch_size]
attn_metadata.spec_sequence_masks[num_spec_decodes:].fill_(False)
assert attn_metadata.non_spec_token_indx is not None
assert attn_metadata.spec_token_indx is not None
non_spec_tokens = attn_metadata.non_spec_token_indx
spec_tokens = attn_metadata.spec_token_indx
self.non_spec_token_indx[: non_spec_tokens.size(0)].copy_(
non_spec_tokens,
non_blocking=True,
)
self.spec_token_indx[: spec_tokens.size(0)].copy_(
spec_tokens,
non_blocking=True,
)
attn_metadata.non_spec_token_indx = self.non_spec_token_indx[: non_spec_tokens.size(0)]
attn_metadata.spec_token_indx = self.spec_token_indx[: spec_tokens.size(0)]
self.spec_query_start_loc[: num_spec_decodes + 1].copy_(
spec_query_start_loc,
non_blocking=True,
)
attn_metadata.spec_query_start_loc = self.spec_query_start_loc[: graph_batch_size + 1]
query_padding = attn_metadata.spec_query_start_loc[num_spec_decodes + 1 :]
if query_padding.numel() > 0:
query_padding.copy_(
spec_query_start_loc[-1].expand_as(query_padding),
non_blocking=True,
)
self.num_accepted_tokens[:num_spec_decodes].copy_(
num_accepted_tokens,
non_blocking=True,
)
attn_metadata.num_accepted_tokens = self.num_accepted_tokens[:graph_batch_size]
attn_metadata.num_accepted_tokens[num_spec_decodes:].fill_(0)
self._attach_spec_decode_metadata(attn_metadata)
def _pad_decode_metadata(
self,
attn_metadata: GDNAttentionMetadata,
graph_batch_size: int,
) -> None:
state_indices = attn_metadata.non_spec_state_indices_tensor
query_start_loc = attn_metadata.non_spec_query_start_loc
assert state_indices is not None
assert query_start_loc is not None
(
attn_metadata.non_spec_state_indices_tensor,
attn_metadata.non_spec_query_start_loc,
) = self._pad_non_spec_decode_graph_inputs(
state_indices,
query_start_loc,
num_decode_tokens=attn_metadata.num_decode_tokens,
graph_batch_size=graph_batch_size,
)
self._attach_non_spec_decode_metadata(
attn_metadata,
attn_metadata.non_spec_state_indices_tensor,
)
def build( # type: ignore[override]
self,
common_prefix_len: int,
common_attn_metadata: CommonAttentionMetadata,
num_accepted_tokens: torch.Tensor | None = None,
num_decode_draft_tokens_cpu: torch.Tensor | None = None,
fast_build: bool = False,
) -> GDNAttentionMetadata:
use_full_graph = self.use_full_cuda_graph
self.use_full_cuda_graph = False
try:
attn_metadata = super().build(
common_prefix_len,
common_attn_metadata,
num_accepted_tokens,
num_decode_draft_tokens_cpu,
fast_build,
)
finally:
self.use_full_cuda_graph = use_full_graph
if not use_full_graph:
return attn_metadata
graph_batch_size = common_attn_metadata.num_reqs
if (
attn_metadata.num_prefills == 0
and attn_metadata.num_decodes == 0
and attn_metadata.num_spec_decodes <= self.decode_cudagraph_max_bs
and attn_metadata.num_spec_decode_tokens <= self.decode_cudagraph_max_bs
):
self._pad_spec_decode_metadata(attn_metadata, graph_batch_size)
elif (
attn_metadata.num_prefills == 0
and attn_metadata.num_spec_decodes == 0
and attn_metadata.num_decodes <= self.decode_cudagraph_max_bs
):
self._pad_decode_metadata(attn_metadata, graph_batch_size)
return attn_metadata
# Keep the name introduced by the 310P ACL graph padding patch so existing
# imports and tests from that patch continue to work after rebasing onto
# upstream/main, whose class name is GDNAttentionMetadataBuilder310.
AscendGDNAttentionMetadataBuilder310 = GDNAttentionMetadataBuilder310
class AscendGDNAttentionBackend310(AscendGDNAttentionBackend):
@staticmethod
def get_builder_cls() -> type[AscendGDNAttentionMetadataBuilder310]:
return AscendGDNAttentionMetadataBuilder310

View File

@@ -0,0 +1,68 @@
import torch
import torch.nn.functional as F
import torch_npu
from vllm.model_executor.layers.layernorm import RMSNormGated
from vllm_ascend.ops.layernorm import AscendGemmaRMSNorm, AscendRMSNorm
class AscendRMSNorm310(AscendRMSNorm):
def forward_oot(
self,
x: torch.Tensor,
residual: torch.Tensor | None = None,
) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
if residual is not None:
x, _, residual = torch_npu.npu_add_rms_norm(x, residual, self.weight, self.variance_epsilon)
if self.bias is not None:
x.add_(self.bias)
return x, residual
x, _ = torch_npu.npu_rms_norm(x, self.weight, self.variance_epsilon)
if self.bias is not None:
x.add_(self.bias)
return x
class AscendGemmaRMSNorm310(AscendGemmaRMSNorm):
def forward_oot(
self,
x: torch.Tensor,
residual: torch.Tensor | None = None,
) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
if residual is not None:
orig_dtype = residual.dtype
x = x + residual.to(x.dtype)
residual = x.to(orig_dtype)
x, _ = torch_npu.npu_rms_norm(x, 1.0 + self.weight, self.variance_epsilon)
return x, residual
x, _ = torch_npu.npu_rms_norm(x, 1.0 + self.weight, self.variance_epsilon)
return x
class AscendRMSNormGated310(RMSNormGated):
def _apply_activation(self, z: torch.Tensor) -> torch.Tensor:
if self.activation == "sigmoid":
return torch.sigmoid(z)
if self.activation in ("silu", "swish"):
return F.silu(z)
raise AssertionError(f"Unsupported activation: {self.activation}")
def forward_oot(
self,
x: torch.Tensor,
z: torch.Tensor | None = None,
) -> torch.Tensor:
if self.group_size is not None:
return super().forward_native(x, z)
if z is not None and not self.norm_before_gate:
x = torch.mul(x, self._apply_activation(z))
x, _ = torch_npu.npu_rms_norm(x, self.weight, self.eps)
if z is not None and self.norm_before_gate:
x = torch.mul(x, self._apply_activation(z))
return x

View File

@@ -0,0 +1,161 @@
#
# 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.
#
import einops
import torch
import torch.nn.functional as F
import torch_npu
from vllm.model_executor.layers.attention.mm_encoder_attention import MMEncoderAttention # type: ignore
MIN_PAD_SIZE: int = 64 # min_size to pad weight
MAX_PAD_SIZE: int = 128 # max_size to pad weight
# Use seq_lens CPU cache to avoid frequent d2h copy.
# AscendMMEncoderAttention310 will copy the cu_seqlens from NPU to CPU in every
# forward, since the op _npu_flash_attention_unpad() requires CPU cu_seqlens
# (otherwise it will break down).
# Thus, we use seq_lens_cpu_cache to cache this tensor, since it's shared
# between all layers, but may change in different forward step. When the
# current layer_index is 0, we update the cache, otherwise we directly use the
# cache to avoid frequent diff and copy operations, which are costful.
seq_lens_cpu_cache: torch.Tensor = None
def is_approximate_calculation_supported() -> bool:
return hasattr(torch_npu, "_npu_flash_attention_unpad_v2")
class AscendMMEncoderAttention310(MMEncoderAttention):
def __init__(
self,
num_heads: int,
head_size: int,
scale: float | None = None,
num_kv_heads: int | None = None,
prefix: str = "",
) -> None:
"""
Args:
num_heads: number of attention heads per partition.
head_size: hidden_size per attention head.
scale: scale factor.
num_kv_heads: number of kv heads.
prefix: This has no effect, it is only here to make it easier to
swap between Attention and MMEncoderAttention.
multimodal_config: configs for multi-modal.
"""
super().__init__(
num_heads=num_heads,
head_size=head_size,
scale=scale,
num_kv_heads=num_kv_heads,
prefix=prefix,
)
self.enable_pad = self.head_size > MIN_PAD_SIZE and self.head_size < MAX_PAD_SIZE
self.scale_value = self.head_size**-0.5
self.support_approximate_calculation = is_approximate_calculation_supported()
def _reshape_qkv_to_3d(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
bsz: int,
q_len: int,
kv_len: int,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""
Reshape query, key, value to 3D tensors:
(batch_size * seq_len, num_heads, head_size)
"""
query = query.view(bsz * q_len, self.num_heads, self.head_size)
key = key.view(bsz * kv_len, self.num_kv_heads, self.head_size)
value = value.view(bsz * kv_len, self.num_kv_heads, self.head_size)
self.num_queries_per_kv = self.num_heads // self.num_kv_heads
if (num_repeat := self.num_queries_per_kv) > 1:
# Handle MQA and GQA
key = torch.repeat_interleave(key, num_repeat, dim=1)
value = torch.repeat_interleave(value, num_repeat, dim=1)
return query, key, value
def forward_oot(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
cu_seqlens: torch.Tensor | None = None,
max_seqlen: torch.Tensor | None = None, # Only used for Flash Attention
sequence_lengths: torch.Tensor | None = None,
):
bsz, q_len = query.size()[:2]
kv_len = key.size(1)
is_reshaped = query.dim() == 4
# Directly use seq_lens cpu cache to avoid d2h copy.
if cu_seqlens is None:
cu_seqlens = torch.arange(0, (bsz + 1) * q_len, step=q_len, dtype=torch.int32, device="cpu")
seq_lens_cpu = torch.diff(cu_seqlens).to("cpu")
# q, k, v: [b, s, head, head_dim] -> [b * s, head, head_dim]
q, k, v = self._reshape_qkv_to_3d(query, key, value, bsz, q_len, kv_len)
if self.enable_pad:
origin_shape = q.shape[-1]
pad_len = MAX_PAD_SIZE - origin_shape
# [b * s, head, head_dim] -> [b * s, head, MAX_PAD_SIZE]
q = F.pad(q, (0, pad_len), mode="constant", value=0)
k = F.pad(k, (0, pad_len), mode="constant", value=0)
v = F.pad(v, (0, pad_len), mode="constant", value=0)
context_layer = torch.empty_like(q)
# TODO: The current torch_npu version 2.10.0 does not support
# _npu_flash_attention_unpad_v2. Once it is supported, drop the
# else branch below and always use the v2 op.
if self.support_approximate_calculation:
torch_npu._npu_flash_attention_unpad_v2(
query=q,
key=k,
value=v,
seq_len=seq_lens_cpu,
scale_value=self.scale_value,
num_heads=self.num_heads,
num_kv_heads=self.num_kv_heads,
out=context_layer,
kernel_type=2,
)
else:
torch_npu._npu_flash_attention_unpad(
query=q,
key=k,
value=v,
seq_len=seq_lens_cpu,
scale_value=self.scale_value,
num_heads=self.num_heads,
num_kv_heads=self.num_kv_heads,
out=context_layer,
)
if self.enable_pad:
context_layer = context_layer[..., :origin_shape]
if is_reshaped:
context_layer = einops.rearrange(context_layer, "(b s) h d -> b s h d", b=bsz).contiguous()
else:
context_layer = einops.rearrange(context_layer, "(b s) h d -> b s (h d)", b=bsz).contiguous()
return context_layer

View File

@@ -0,0 +1,35 @@
#
# 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.
#
import torch
# 310P RC: non_blocking H2D copy in rot_pos_emb can race with subsequent indexing.
def rot_pos_emb_310(self, grid_thw: list[list[int]]):
max_grid_size = max(max(h, w) for _, h, w in grid_thw)
pos_ids = [
self.rot_pos_ids(h, w, self.spatial_merge_size)
if t == 1
else self.rot_pos_ids(h, w, self.spatial_merge_size).repeat(t, 1)
for t, h, w in grid_thw
]
pos_ids = torch.cat(pos_ids, dim=0).to(self.device, non_blocking=False)
cos, sin = self.rotary_pos_emb.get_cos_sin(max_grid_size)
cos_combined = cos[pos_ids].flatten(1)
sin_combined = sin[pos_ids].flatten(1)
return cos_combined, sin_combined

View File

@@ -0,0 +1,273 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
#
# 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.
# This file is a part of the vllm-ascend project.
#
from __future__ import annotations
from typing import Any
import torch
import torch_npu
from vllm.model_executor.layers.rotary_embedding import MRotaryEmbedding
from vllm.model_executor.layers.rotary_embedding.common import ApplyRotaryEmb
from vllm.model_executor.layers.rotary_embedding.mrope import apply_interleaved_rope
from vllm_ascend.ops.rotary_embedding import AscendRotaryEmbedding, get_cos_and_sin_slice, update_cos_sin
# Filled once per model forward in NPUModelRunner310._model_forward; read by every MRoPE layer.
_mrope_cos_slice: torch.Tensor | None = None
_mrope_sin_slice: torch.Tensor | None = None
def _apply_rotary_mrope_torch(
q_rot: torch.Tensor,
k_rot: torch.Tensor,
cos: torch.Tensor,
sin: torch.Tensor,
is_neox_style: bool,
) -> tuple[torch.Tensor, torch.Tensor]:
"""PyTorch path aligned with vLLM MRotaryEmbedding.forward_native -> ApplyRotaryEmb."""
half = cos.shape[-1] // 2
cos_h = cos[0, :, 0, :half].contiguous()
sin_h = sin[0, :, 0, :half].contiguous()
q_out = ApplyRotaryEmb.forward_static(q_rot[0], cos_h, sin_h, is_neox_style)
k_out = ApplyRotaryEmb.forward_static(k_rot[0], cos_h, sin_h, is_neox_style)
return q_out.unsqueeze(0), k_out.unsqueeze(0)
def merge_mrope_cos_sin_for_apply(
cos: torch.Tensor,
sin: torch.Tensor,
mrope_section: list[int],
mrope_interleaved: bool,
) -> tuple[torch.Tensor, torch.Tensor]:
if mrope_interleaved:
return (
apply_interleaved_rope(cos, mrope_section),
apply_interleaved_rope(sin, mrope_section),
)
return (
torch.cat([m[i] for i, m in enumerate(cos.split(mrope_section, dim=-1))], dim=-1),
torch.cat([m[i] for i, m in enumerate(sin.split(mrope_section, dim=-1))], dim=-1),
)
def set_mrope_apply_rotary_slices(
cos_sin_cache: torch.Tensor,
positions: torch.Tensor,
*,
mrope_section: list[int] | None = None,
mrope_interleaved: bool = False,
capacity_tokens: int = 0,
) -> None:
"""Build cos/sin views for `npu_apply_rotary_pos_emb` from positions; must run once per forward before layers."""
global _mrope_cos_slice
global _mrope_sin_slice
assert positions.ndim in (1, 2), "M-RoPE positions must be [num_tokens] or [3, num_tokens]."
cos_sin = cos_sin_cache[positions]
cos, sin = cos_sin.chunk(2, dim=-1)
if positions.ndim == 2:
assert positions.shape[0] == 3, "MRoPE expects positions [3, num_tokens] (T/H/W)."
assert mrope_section is not None
cos, sin = merge_mrope_cos_sin_for_apply(
cos,
sin,
list(mrope_section),
mrope_interleaved,
)
# `npu_apply_rotary_pos_emb` follows ApplyRotaryPosEmbV2 semantics:
# q_embed = q * cos + rotate(q) * sin, where cos/sin have full rotary dim.
# MRoPE merge above gives half-dim cos/sin, so expand to full dim here.
cos = torch.cat((cos, cos), dim=-1)
sin = torch.cat((sin, sin), dim=-1)
num_tokens = positions.shape[-1]
cos_view = cos.contiguous().view(1, num_tokens, 1, -1)
sin_view = sin.contiguous().view(1, num_tokens, 1, -1)
# Keep stable storage across forwards for graph replay.
if _mrope_cos_slice is None or _mrope_sin_slice is None:
capacity = capacity_tokens if capacity_tokens is not None else num_tokens
if capacity < num_tokens:
capacity = num_tokens
_mrope_cos_slice = torch.empty(
(1, capacity, 1, cos_view.shape[-1]),
dtype=cos_view.dtype,
device=cos_view.device,
)
_mrope_sin_slice = torch.empty(
(1, capacity, 1, sin_view.shape[-1]),
dtype=sin_view.dtype,
device=sin_view.device,
)
_mrope_cos_slice[:, :num_tokens].copy_(cos_view)
_mrope_sin_slice[:, :num_tokens].copy_(sin_view)
def _rope_forward_oot(
self,
positions: torch.Tensor,
query: torch.Tensor,
key: torch.Tensor,
is_neox_style: bool,
offsets: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
query_shape, key_shape = query.shape, key.shape
if self.cos_sin_cache.device != query.device:
self.cos_sin_cache = self.cos_sin_cache.to(query.device)
if self.cos_sin_cache.dtype != query.dtype:
self.cos_sin_cache = self.cos_sin_cache.to(query.dtype)
# This flag should set to True when doing drafting.
if getattr(self, "_is_drafting_update_enabled", False):
update_cos_sin(positions)
cos, sin = get_cos_and_sin_slice()
if offsets is not None:
raise NotImplementedError("Batched rotary embedding is currently not supported on NPU.")
rotary_mode = "half" if is_neox_style else "interleave"
if self.head_size == 128 and self.cos_sin_cache.shape[-1] == 128:
query = query.contiguous().view(1, query.shape[0], -1, self.head_size)
key = key.contiguous().view(1, key.shape[0], -1, self.head_size)
query, key = torch_npu.npu_apply_rotary_pos_emb(query, key, cos, sin, rotary_mode=rotary_mode)
elif self.rotary_dim < self.head_size:
num_tokens = query.shape[0]
query = query.view(num_tokens, -1, self.head_size)
key = key.view(num_tokens, -1, self.head_size)
q_rot = query[..., : self.rotary_dim]
q_pass = query[..., self.rotary_dim :]
k_rot = key[..., : self.rotary_dim]
k_pass = key[..., self.rotary_dim :]
if self.rotary_dim == 64:
q_rot = q_rot.contiguous().view(1, num_tokens, -1, self.rotary_dim)
k_rot = k_rot.contiguous().view(1, num_tokens, -1, self.rotary_dim)
q_rot, k_rot = torch_npu.npu_apply_rotary_pos_emb(q_rot, k_rot, cos, sin, rotary_mode=rotary_mode)
else:
q_rot = q_rot.contiguous().view(num_tokens, -1)
k_rot = k_rot.contiguous().view(num_tokens, -1)
torch_npu._npu_rotary_embedding(
positions,
q_rot,
k_rot,
self.rotary_dim,
self.cos_sin_cache,
is_neox_style,
)
q_rot = q_rot.view(num_tokens, -1, self.rotary_dim)
k_rot = k_rot.view(num_tokens, -1, self.rotary_dim)
query = torch.cat((q_rot, q_pass), dim=-1).reshape(query_shape)
key = torch.cat((k_rot, k_pass), dim=-1).reshape(key_shape)
else:
query = query.contiguous().view(query.shape[0], -1)
key = key.contiguous().view(key.shape[0], -1)
torch_npu._npu_rotary_embedding(
positions,
query,
key,
self.head_size,
self.cos_sin_cache,
is_neox_style,
)
return query.view(query_shape), key.view(key_shape)
class AscendMRotaryEmbedding310(MRotaryEmbedding):
def forward_oot(
self,
positions: torch.Tensor,
query: torch.Tensor,
key: torch.Tensor,
):
query_shape, key_shape = query.shape, key.shape
# MRoPE T/H/W layout is handled in `merge_mrope_cos_sin_for_apply` (mrope_interleaved).
# Here `rotary_mode` matches vLLM ApplyRotaryEmb: half = neox chunk, interleave = GPT-J pairs.
rotary_mode = "half" if self.is_neox_style else "interleave"
num_tokens = query.shape[0]
if _mrope_cos_slice is None or _mrope_sin_slice is None:
raise RuntimeError(
"MRoPE cos/sin slices are not initialized. Call set_mrope_apply_rotary_slices before forward."
)
cos, sin = _mrope_cos_slice[:, :num_tokens], _mrope_sin_slice[:, :num_tokens]
is_partial_rope = self.rotary_dim < self.head_size
if is_partial_rope:
query = query.view(num_tokens, -1, self.head_size)
key = key.view(num_tokens, -1, self.head_size)
q_pass = query[..., self.rotary_dim :]
k_pass = key[..., self.rotary_dim :]
q_rot = query[..., : self.rotary_dim].contiguous().view(1, num_tokens, -1, self.rotary_dim)
k_rot = key[..., : self.rotary_dim].contiguous().view(1, num_tokens, -1, self.rotary_dim)
else:
q_rot = query.contiguous().view(1, num_tokens, -1, self.head_size)
k_rot = key.contiguous().view(1, num_tokens, -1, self.head_size)
# `npu_apply_rotary_pos_emb` only supports rotary_dim 64 or 128.
use_npu_apply = self.rotary_dim in (64, 128)
if use_npu_apply:
q_rot, k_rot = torch_npu.npu_apply_rotary_pos_emb(q_rot, k_rot, cos, sin, rotary_mode=rotary_mode)
else:
q_rot, k_rot = _apply_rotary_mrope_torch(q_rot, k_rot, cos, sin, self.is_neox_style)
if is_partial_rope:
q_rot = q_rot.view(num_tokens, -1, self.rotary_dim)
k_rot = k_rot.view(num_tokens, -1, self.rotary_dim)
query = torch.cat((q_rot, q_pass), dim=-1).reshape(query_shape)
key = torch.cat((k_rot, k_pass), dim=-1).reshape(key_shape)
else:
query = q_rot.view(query_shape)
key = k_rot.view(key_shape)
return query, key
def prepare_mrope_cos_sin_slices_from_runner(runner: Any, positions: torch.Tensor) -> None:
"""Resolve MRoPE embedding from the runner and populate `_mrope_cos_slice` / `_mrope_sin_slice`."""
emb = getattr(runner, "_mrope_embedding", None)
if emb is None:
emb = next(module for module in runner.model.modules() if isinstance(module, AscendMRotaryEmbedding310))
runner._mrope_embedding = emb
assert isinstance(emb, AscendMRotaryEmbedding310)
set_mrope_apply_rotary_slices(
emb.cos_sin_cache,
positions,
mrope_section=emb.mrope_section,
mrope_interleaved=emb.mrope_interleaved,
capacity_tokens=runner.max_num_tokens,
)
class AscendRotaryEmbedding310(AscendRotaryEmbedding):
_is_drafting_update_enabled: bool = False
@classmethod
def set_rope_position_flag_310p(cls, state: bool):
cls._is_drafting_update_enabled = state
def forward_oot(
self,
positions: torch.Tensor,
query: torch.Tensor,
key: torch.Tensor,
offsets: torch.Tensor | None = None,
is_neox_style_override: bool | None = None,
):
is_neox_style = self.is_neox_style
if is_neox_style_override is not None:
is_neox_style = is_neox_style_override
return _rope_forward_oot(self, positions, query, key, is_neox_style, offsets)

View File

@@ -0,0 +1,82 @@
#
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
#
# 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 __future__ import annotations
import torch
import torch.nn.functional as F
from vllm.model_executor.layers.quantization.base_config import QuantizationConfig
from vllm.model_executor.layers.vocab_parallel_embedding import (
DEFAULT_VOCAB_PADDING_SIZE,
UnquantizedEmbeddingMethod,
)
from vllm_ascend.ops.vocab_parallel_embedding import AscendParallelLMHead, AscendVocabParallelEmbedding
from vllm_ascend.utils import maybe_trans_nz
class AscendUnquantizedEmbeddingMethod310(UnquantizedEmbeddingMethod):
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
layer.weight_nz = maybe_trans_nz(layer.weight)
def apply(
self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: torch.Tensor | None = None,
) -> torch.Tensor:
return F.linear(x, layer.weight_nz, bias)
class AscendVocabParallelEmbedding310(AscendVocabParallelEmbedding):
def __init__(
self,
num_embeddings: int,
embedding_dim: int,
params_dtype: torch.dtype | None = None,
org_num_embeddings: int | None = None,
padding_size: int = DEFAULT_VOCAB_PADDING_SIZE,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
):
super().__init__(
num_embeddings, embedding_dim, params_dtype, org_num_embeddings, padding_size, quant_config, prefix
)
if quant_config is None:
self.quant_method = AscendUnquantizedEmbeddingMethod310()
class AscendParallelLMHead310(AscendParallelLMHead):
"""
Register ParallelLMHead as a custom op for Atlas 310p.
"""
def __init__(
self,
num_embeddings: int,
embedding_dim: int,
bias: bool = False,
params_dtype: torch.dtype | None = None,
org_num_embeddings: int | None = None,
padding_size: int = DEFAULT_VOCAB_PADDING_SIZE,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
):
super().__init__(
num_embeddings, embedding_dim, bias, params_dtype, org_num_embeddings, padding_size, quant_config, prefix
)
if quant_config is None:
self.quant_method = AscendUnquantizedEmbeddingMethod310()