11
vllm_ascend/_310p/ops/fla/__init__.py
Normal file
11
vllm_ascend/_310p/ops/fla/__init__.py
Normal 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",
|
||||
]
|
||||
586
vllm_ascend/_310p/ops/fla/chunk_gated_delta_rule.py
Normal file
586
vllm_ascend/_310p/ops/fla/chunk_gated_delta_rule.py
Normal 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
|
||||
@@ -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
|
||||
62
vllm_ascend/_310p/ops/fla/fused_gdn_gating.py
Normal file
62
vllm_ascend/_310p/ops/fla/fused_gdn_gating.py
Normal 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
|
||||
227
vllm_ascend/_310p/ops/fla/fused_recurrent_gated_delta_rule.py
Normal file
227
vllm_ascend/_310p/ops/fla/fused_recurrent_gated_delta_rule.py
Normal 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
|
||||
429
vllm_ascend/_310p/ops/fla/gdn_310.py
Normal file
429
vllm_ascend/_310p/ops/fla/gdn_310.py
Normal 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("", [])
|
||||
42
vllm_ascend/_310p/ops/fla/idex.py
Normal file
42
vllm_ascend/_310p/ops/fla/idex.py
Normal 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)
|
||||
43
vllm_ascend/_310p/ops/fla/l2norm.py
Normal file
43
vllm_ascend/_310p/ops/fla/l2norm.py
Normal 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)
|
||||
Reference in New Issue
Block a user