[NPU]Support GLM5.2 and FP8 DSA&Indexer kvcache for 950 (#38250)

This commit is contained in:
LinyuanLi
2026-09-10 09:10:10 +08:00
committed by GitHub
parent a7e00b7576
commit 3ff226ba8f
10 changed files with 566 additions and 136 deletions
@@ -314,6 +314,7 @@ class AscendAttnBackend(AttentionBackend):
)
self.page_size = model_runner.page_size
self.model_dtype = model_runner.model_config.dtype
self.kv_cache_dtype = model_runner.kv_cache_dtype
self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA
if self.use_mla:
self.kv_lora_rank = model_runner.model_config.kv_lora_rank
@@ -1178,28 +1179,58 @@ class AscendAttnBackend(AttentionBackend):
if topk_indices is not None:
topk_indices = self._pad_topk_indices(topk_indices, q_nope.shape[0])
topk_indices = _expand_dsa_sparse_indices(topk_indices)
attn_out, _, _ = torch_npu.npu_sparse_flash_attention(
query=q_nope,
key=k_nope,
value=k_nope,
query_rope=q_pe,
key_rope=k_pe,
sparse_indices=topk_indices,
scale_value=layer.scaling,
actual_seq_lengths_query=actual_seq_qlen.to(
device=q_nope.device, dtype=torch.int32
),
actual_seq_lengths_kv=actual_seq_lengths_kv.to(
device=q_nope.device, dtype=torch.int32
),
block_table=self.forward_metadata.block_tables,
sparse_block_size=1,
layout_query="TND",
layout_kv="PA_BSND",
sparse_mode=3,
attention_mode=2,
return_softmax_lse=False,
)
if self.kv_cache_dtype == torch.float8_e4m3fn:
assert q_nope.dtype == q_pe.dtype == torch.bfloat16
packed = k_nope.view(torch.float8_e4m3fn)
attn_out = torch_npu.npu_kv_quant_sparse_flash_attention(
query=torch.cat((q_nope, q_pe), dim=-1).contiguous(),
key=packed,
value=packed,
sparse_indices=topk_indices,
scale_value=layer.scaling,
key_quant_mode=2,
value_quant_mode=2,
key_dequant_scale=None,
value_dequant_scale=None,
actual_seq_lengths_query=actual_seq_qlen.to(
device=q_nope.device, dtype=torch.int32
),
actual_seq_lengths_kv=actual_seq_lengths_kv.to(
device=q_nope.device, dtype=torch.int32
),
block_table=self.forward_metadata.block_tables,
sparse_block_size=1,
layout_query="TND",
layout_kv="PA_BSND",
sparse_mode=3,
attention_mode=2,
quant_scale_repo_mode=1,
tile_size=128,
rope_head_dim=self.qk_rope_head_dim,
)
else:
attn_out, _, _ = torch_npu.npu_sparse_flash_attention(
query=q_nope,
key=k_nope,
value=k_nope,
query_rope=q_pe,
key_rope=k_pe,
sparse_indices=topk_indices,
scale_value=layer.scaling,
actual_seq_lengths_query=actual_seq_qlen.to(
device=q_nope.device, dtype=torch.int32
),
actual_seq_lengths_kv=actual_seq_lengths_kv.to(
device=q_nope.device, dtype=torch.int32
),
block_table=self.forward_metadata.block_tables,
sparse_block_size=1,
layout_query="TND",
layout_kv="PA_BSND",
sparse_mode=3,
attention_mode=2,
return_softmax_lse=False,
)
return attn_out
@@ -1219,8 +1250,10 @@ class AscendAttnBackend(AttentionBackend):
slopes: Optional[torch.Tensor] = None,
):
if is_mla_preprocess_enabled() and self.use_mla:
# MLAPO and MLAPROLOG do save kv_cache
save_kv_cache = False
# DSA callers set save_kv_cache based on whether preprocessing was used.
# Only override it for the existing non-sparse MLA path.
if topk_indices is None:
save_kv_cache = False
if self.is_dllm_model:
return self.forward_dllm(
q,
@@ -2600,8 +2633,10 @@ class AscendAttnBackend(AttentionBackend):
**kwargs,
):
if is_mla_preprocess_enabled() and self.use_mla:
# MLAPO does saving kv_cache
save_kv_cache = False
# DSA callers set save_kv_cache based on whether preprocessing was used.
# Only override it for the existing non-sparse MLA path.
if topk_indices is None:
save_kv_cache = False
if topk_indices is not None:
if self.enable_sparsity_driven_kv_offload:
from sglang.srt.hardware_backend.npu.sparsity_driven_kv_offload.attention import (
@@ -5,7 +5,7 @@ from typing import TYPE_CHECKING, Optional
import torch
import torch.nn.functional as F
from sglang.srt.hardware_backend.npu.utils import npu_format_cast
from sglang.srt.hardware_backend.npu.utils import is_npu_arch35, npu_format_cast
from sglang.srt.model_executor.forward_context import (
get_attn_backend,
get_token_to_kv_pool,
@@ -99,9 +99,11 @@ class NPUFusedMLAPreprocess(torch.nn.Module):
self.qk_rope_head_dim = qk_rope_head_dim # 64
self.qk_head_dim = qk_nope_head_dim + qk_rope_head_dim
self.v_head_dim = v_head_dim
self.q_b_proj_weight_scale = self.q_b_proj.weight_scale.view(1, -1).to(
torch.float
q_b_scale = getattr(self.q_b_proj, "weight_scale", None)
self.q_b_proj_weight_scale = (
q_b_scale.view(1, -1).to(torch.float) if q_b_scale is not None else None
)
self.is_npu_arch35 = is_npu_arch35()
def preprocess_weights(self, hidden_states):
self.dummy = torch.zeros(
@@ -241,13 +243,90 @@ class NPUFusedMLAPreprocess(torch.nn.Module):
)
def mlaprolog_preprocess_weight(self):
self.qkv_a_proj.weight.data = self.qkv_a_proj.weight.data.transpose(0, 1)
qkv_a_proj_weight_q = self.qkv_a_proj.weight.data[:, : self.q_lora_rank].clone()
qkv_a_proj_weight_kv = self.qkv_a_proj.weight.data[
:, self.q_lora_rank :
].clone()
self.q_a_proj_weight = npu_format_cast(qkv_a_proj_weight_q)
self.kv_a_proj_weight = npu_format_cast(qkv_a_proj_weight_kv)
# MLAPrologV3 weight quantization modes (self.weight_quant_mode) used here:
# 0: No weight quantization. QKV-A and Q-B weights are FP16/BF16.
# 1: Partial INT8 quantization. Only weight_uq_qr (Q-B projection)
# is INT8; weight_dq and weight_dkv_kr (QKV-A projection) remain
# FP16/BF16. dequant_scale_w_uq_qr is required.
# 3: MXFP8 quantization. token_x, weight_dq, weight_uq_qr, and
# weight_dkv_kr use MXFP8 with their corresponding dequant scales.
# weight_uk remains unquantized.
from sglang.srt.hardware_backend.npu.quantization.linear_method_npu import (
NPUMXFP8LinearMethod,
)
projections = (self.qkv_a_proj, self.q_b_proj)
kernels = [
getattr(getattr(layer, "scheme", None), "kernel", layer.quant_method)
for layer in projections
]
is_mxfp8 = [isinstance(kernel, NPUMXFP8LinearMethod) for kernel in kernels]
if any(is_mxfp8):
if not all(is_mxfp8):
raise RuntimeError(
"MLAProlog MXFP8 requires both QKV-A and Q-B to use MXFP8"
)
expected_shapes = (
(
self.qkv_a_proj.input_size,
self.q_lora_rank + self.kv_lora_rank + self.qk_rope_head_dim,
),
(self.q_lora_rank, self.num_local_heads * self.qk_head_dim),
)
checkpoint_scales = []
for layer, (k_dim, n_dim) in zip(projections, expected_shapes):
scale = getattr(layer, "weight_scale_inv", None)
if (
layer.weight.dtype != torch.float8_e4m3fn
or tuple(layer.weight.shape) != (k_dim, n_dim)
or k_dim % 64
or scale is None
or tuple(scale.shape) != (k_dim // 64, n_dim, 2)
or scale.dtype not in (torch.uint8, torch.float8_e8m0fnu)
):
raise RuntimeError(
"MLAProlog requires the NPUMXFP8 ND weight/paired-scale layout"
)
# Invert the mainline post-load views. No source copy or live mutation.
checkpoint_scales.append(
scale.data.transpose(0, 1).reshape(n_dim, k_dim // 32)
)
qkv_weight = self.qkv_a_proj.weight.data
qkv_scale, qb_scale = checkpoint_scales
self.qkv_a_proj_scale_q = (
qkv_scale[: self.q_lora_rank].contiguous().view(torch.float8_e8m0fnu)
)
self.qkv_a_proj_scale_kv = (
qkv_scale[self.q_lora_rank :].contiguous().view(torch.float8_e8m0fnu)
)
self.q_b_proj_scale = qb_scale.contiguous().view(torch.float8_e8m0fnu)
self.q_b_proj_weight = npu_format_cast(
self.q_b_proj.weight.data.contiguous()
)
self.weight_quant_mode = 3
else:
if self.qkv_a_proj.weight.dtype not in (torch.float16, torch.bfloat16):
raise RuntimeError("Unsupported MLAProlog QKV-A weight format")
qkv_weight = self.qkv_a_proj.weight.data.transpose(0, 1)
if self.q_b_proj.weight.dtype in (torch.float16, torch.bfloat16):
self.weight_quant_mode = 0
self.q_b_proj_weight = npu_format_cast(
self.q_b_proj.weight.data.transpose(0, 1).contiguous()
)
elif (
self.q_b_proj.weight.dtype == torch.int8
and self.q_b_proj_weight_scale is not None
):
self.weight_quant_mode = 1
self.q_b_proj_weight = self.q_b_proj.weight
else:
raise RuntimeError("Unsupported MLAProlog Q-B weight format")
self.q_a_proj_weight = npu_format_cast(
qkv_weight[:, : self.q_lora_rank].contiguous()
)
self.kv_a_proj_weight = npu_format_cast(
qkv_weight[:, self.q_lora_rank :].contiguous()
)
def get_sin_cos(self, positions):
cos_sin = self.rotary_emb.cos_sin_cache[positions]
@@ -434,10 +513,27 @@ class NPUFusedMLAPreprocess(torch.nn.Module):
self.has_preprocess_weights = True
self.cos, self.sin = self.get_sin_cos(positions)
k_cache, v_cache, slot_mapping = self.get_kv_cache_and_cache_idx(forward_batch)
pool = get_token_to_kv_pool()
packed = pool.dsa_kv_cache_store_fp8
if packed and self.weight_quant_mode != 3:
raise RuntimeError(
"Packed FP8 KV with MLAProlog requires MXFP8 QKV-A and Q-B weights; "
"use BF16 draft KV for BF16 draft weights"
)
token_x = hidden_states
if self.weight_quant_mode == 3:
token_x, token_x_scale = torch.ops.npu.npu_dynamic_mx_quant(
hidden_states.reshape(-1, hidden_states.shape[-1]).contiguous(),
axis=1,
dst_type=torch.float8_e4m3fn,
block_size=32,
scale_alg=None,
)
token_x_scale = token_x_scale.contiguous().reshape(token_x.shape[0], -1)
mla_prolog_input_args = {
"token_x": hidden_states,
"token_x": token_x,
"weight_dq": self.q_a_proj_weight,
"weight_uq_qr": self.q_b_proj.weight,
"weight_uq_qr": self.q_b_proj_weight,
"weight_uk": self.w_kc,
"weight_dkv_kr": self.kv_a_proj_weight,
"rmsnorm_gamma_cq": self.q_a_layernorm.weight,
@@ -447,17 +543,42 @@ class NPUFusedMLAPreprocess(torch.nn.Module):
"kv_cache": k_cache,
"kr_cache": v_cache,
"cache_index": slot_mapping.to(dtype=torch.int64),
"dequant_scale_w_uq_qr": self.q_b_proj_weight_scale,
"rmsnorm_epsilon_cq": self.q_a_layernorm.variance_epsilon,
"rmsnorm_epsilon_ckv": self.kv_a_layernorm.variance_epsilon,
"cache_mode": "PA_BSND",
"cache_mode": "PA_BSND" if packed or not is_fia_nz() else "PA_NZ",
"query_norm_flag": True,
"weight_quant_mode": 1, # 0:no quant; 1:uq_qr: quant; 2: weight_dq,weight_uq_qr,weight_dkv_kr: quant
"weight_quant_mode": self.weight_quant_mode,
}
if self.is_npu_arch35 and pool.index_head_dim is not None:
mla_prolog_input_args.update(
kv_cache_quant_mode=3 if packed else 0,
query_quant_mode=0,
)
if self.weight_quant_mode == 3:
mla_prolog_input_args.update(
dequant_scale_w_dq=self.qkv_a_proj_scale_q,
dequant_scale_w_dkv_kr=self.qkv_a_proj_scale_kv,
dequant_scale_w_uq_qr=self.q_b_proj_scale,
dequant_scale_x=token_x_scale.view(torch.float8_e8m0fnu),
kc_scale=1.0,
qc_qr_scale=1.0,
quant_scale_ckv=None,
)
elif self.weight_quant_mode == 1:
mla_prolog_input_args["dequant_scale_w_uq_qr"] = self.q_b_proj_weight_scale
if packed:
mla_prolog_input_args.update(
ckvkr_repo_mode=1, quant_scale_repo_mode=1, tile_size=128
)
import torch_npu
q_nope, q_pe, dequant_scale_q_nope, qr, dequant_q_norm = (
torch.ops.custom.npu_mla_prolog_v3(**mla_prolog_input_args)
torch_npu.npu_mla_prolog_v3(**mla_prolog_input_args)
)
dequant_q_norm = dequant_q_norm.view(hidden_states.shape[0])
if self.weight_quant_mode == 0:
dequant_q_norm = None
elif self.weight_quant_mode == 1:
dequant_q_norm = dequant_q_norm.view(hidden_states.shape[0])
return (
q_pe,
v_cache,
@@ -477,11 +598,15 @@ class NPUFusedMLAPreprocess(torch.nn.Module):
and self.qkv_a_proj.quant_method.quantization_config.get_name()
== "modelslim"
)
# with the mlaprolog enabled, the kv_b_proj layers are unquantized
_is_mlaprolog = hasattr(self.quant_config, "ignore") and any(
re.fullmatch(r".*kv_b_proj", l) for l in self.quant_config.ignore
_is_arch35_dsa = (
self.is_npu_arch35 and get_token_to_kv_pool().index_head_dim is not None
)
if _is_w8a8:
# with the mlaprolog enabled, the kv_b_proj layers are unquantized
_is_mlaprolog = _is_arch35_dsa or (
hasattr(self.quant_config, "ignore")
and any(re.fullmatch(r".*kv_b_proj", l) for l in self.quant_config.ignore)
)
if _is_w8a8 and not _is_arch35_dsa:
return self.forward_mlapo(
positions, hidden_states, forward_batch, zero_allocator
)
@@ -1,4 +1,4 @@
from typing import TYPE_CHECKING, Optional
from typing import TYPE_CHECKING, Optional, Sequence
import torch
@@ -537,6 +537,8 @@ class NPUMLATokenToKVPool(MLATokenToKVPool):
index_head_dim: Optional[int] = None,
start_layer: Optional[int] = None,
end_layer: Optional[int] = None,
indexer_layer_ids: Optional[Sequence[int]] = None,
kv_cache_dim: Optional[int] = None,
):
super(MLATokenToKVPool, self).__init__(
size=size,
@@ -558,6 +560,42 @@ class NPUMLATokenToKVPool(MLATokenToKVPool):
if self.enable_sparsity_driven_kv_offload and self.index_head_dim is None:
raise ValueError("Sparsity-driven KV offload requires an index KV cache.")
if index_head_dim is None:
self.indexer_layer_ids = ()
elif indexer_layer_ids is None:
self.indexer_layer_ids = tuple(
range(self.start_layer, self.start_layer + self.layer_num)
)
else:
self.indexer_layer_ids = tuple(indexer_layer_ids)
self.num_indexer_layers = len(self.indexer_layer_ids)
self.indexer_layer_id_to_slot = {
layer_id: slot for slot, layer_id in enumerate(self.indexer_layer_ids)
}
assert len(self.indexer_layer_id_to_slot) == self.num_indexer_layers
assert all(
self.start_layer <= i < self.start_layer + self.layer_num
for i in self.indexer_layer_ids
)
requested_kv_cache_dim = kv_cache_dim
self.dsa_kv_cache_store_fp8 = (
index_head_dim is not None
and dtype == torch.float8_e4m3fn
and requested_kv_cache_dim is not None
)
if self.dsa_kv_cache_store_fp8:
assert index_head_dim == 128 and kv_lora_rank % 128 == 0
assert requested_kv_cache_dim == (
kv_lora_rank + kv_lora_rank // 128 * 4 + qk_rope_head_dim * 2
)
self.store_dtype = dtype
self.kv_cache_dim = (
requested_kv_cache_dim if self.dsa_kv_cache_store_fp8 else kv_lora_rank
)
self.kr_cache_dim = 0 if self.dsa_kv_cache_store_fp8 else qk_rope_head_dim
self.index_k_scale_buffer = None
self.indexer_hadamard_128 = None
self.custom_mem_pool = None
with self.memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE):
@@ -572,7 +610,7 @@ class NPUMLATokenToKVPool(MLATokenToKVPool):
self.size // self.page_size + 1,
self.page_size,
1,
self.kv_lora_rank,
self.kv_cache_dim,
),
dtype=self.store_dtype,
device=self.device,
@@ -583,16 +621,20 @@ class NPUMLATokenToKVPool(MLATokenToKVPool):
self.size // self.page_size + 1,
self.page_size,
1,
self.qk_rope_head_dim,
self.kr_cache_dim,
),
dtype=(
torch.bfloat16
if self.dsa_kv_cache_store_fp8
else self.store_dtype
),
dtype=self.store_dtype,
device=self.device,
)
self.index_k_buffer = None
if self.index_head_dim is not None:
self.index_k_buffer = torch.zeros(
(
layer_num,
self.num_indexer_layers,
self.size // self.page_size + 1,
self.page_size,
1,
@@ -601,6 +643,19 @@ class NPUMLATokenToKVPool(MLATokenToKVPool):
dtype=self.store_dtype,
device=self.device,
)
if self.dsa_kv_cache_store_fp8 and self.num_indexer_layers > 0:
from sglang.srt.layers.attention.dsa.dsa_npu_indexer import (
create_npu_hadamard_128,
)
self.index_k_scale_buffer = torch.zeros(
(*self.index_k_buffer.shape[:-2], 1),
dtype=torch.float32,
device=self.device,
)
self.indexer_hadamard_128 = create_npu_hadamard_128(
self.index_head_dim, self.device
)
self._finalize_allocation_log(size)
@@ -616,6 +671,8 @@ class NPUMLATokenToKVPool(MLATokenToKVPool):
if getattr(self, "index_k_buffer", None) is not None:
for index_k_cache in self.index_k_buffer:
kv_size_bytes += get_tensor_size_bytes(index_k_cache)
if self.index_k_scale_buffer is not None:
kv_size_bytes += get_tensor_size_bytes(self.index_k_scale_buffer)
return kv_size_bytes
def _raise_if_native_kv_cache_disabled(self):
@@ -641,9 +698,12 @@ class NPUMLATokenToKVPool(MLATokenToKVPool):
def get_state_buf_infos(self):
if self.index_head_dim is None:
return [], [], []
data_ptrs = [self.index_k_buffer[i].data_ptr() for i in range(self.layer_num)]
data_lens = [self.index_k_buffer[i].nbytes for i in range(self.layer_num)]
item_lens = [self.index_k_buffer[i][0].nbytes for i in range(self.layer_num)]
buffers = list(self.index_k_buffer)
if self.index_k_scale_buffer is not None:
buffers += list(self.index_k_scale_buffer)
data_ptrs = [buf.data_ptr() for buf in buffers]
data_lens = [buf.nbytes for buf in buffers]
item_lens = [buf[0].nbytes for buf in buffers]
return data_ptrs, data_lens, item_lens
def get_key_buffer(self, layer_id: int):
@@ -671,8 +731,25 @@ class NPUMLATokenToKVPool(MLATokenToKVPool):
raise RuntimeError("NPU MLA index KV cache is not allocated.")
if self.store_dtype != self.dtype:
return self.index_k_buffer[layer_id - self.start_layer].view(self.dtype)
return self.index_k_buffer[layer_id - self.start_layer]
return self.index_k_buffer[self._get_indexer_slot(layer_id)].view(
self.dtype
)
return self.index_k_buffer[self._get_indexer_slot(layer_id)]
def _get_indexer_slot(self, layer_id: int) -> int:
return self.indexer_layer_id_to_slot[layer_id]
def get_index_k_scale_buffer(self, layer_id: int):
if self.layer_transfer_counter is not None:
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
return self.index_k_scale_buffer[self._get_indexer_slot(layer_id)]
def set_index_k_scale_buffer(self, layer_id: int, loc, scale):
torch_npu.npu_scatter_nd_update_(
self.index_k_scale_buffer[self._get_indexer_slot(layer_id)].view(-1, 1),
loc.view(-1, 1),
scale.view(-1, 1),
)
# for disagg
def get_contiguous_buf_infos(self):
@@ -688,17 +765,46 @@ class NPUMLATokenToKVPool(MLATokenToKVPool):
self.v_buffer[i][0].nbytes for i in range(self.layer_num)
]
if self.index_head_dim is not None:
kv_data_ptrs += [
self.index_k_buffer[i].data_ptr() for i in range(self.layer_num)
]
kv_data_lens += [
self.index_k_buffer[i].nbytes for i in range(self.layer_num)
]
kv_item_lens += [
self.index_k_buffer[i][0].nbytes for i in range(self.layer_num)
]
ptrs, lens, item_lens = self.get_state_buf_infos()
kv_data_ptrs += ptrs
kv_data_lens += lens
kv_item_lens += item_lens
return kv_data_ptrs, kv_data_lens, kv_item_lens
def get_kv_layer_ids(self):
return (
list(range(self.start_layer, self.start_layer + self.layer_num)) * 2
+ self.get_state_layer_ids()
)
def get_state_layer_ids(self):
return list(self.indexer_layer_ids) * (
2 if self.index_k_scale_buffer is not None else 1
)
def _pack_dsa_fp8_kv_cache(self, cache_k, cache_v):
latent = cache_k.reshape(-1, self.kv_lora_rank)
quantized, scale = torch_npu.npu_dynamic_quant(
latent.reshape(-1, 128), dst_type=self.dtype
)
rows = latent.shape[0]
# Opaque record: latent FP8 | rope BF16 bytes | per-tile FP32 scales.
packed = torch.cat(
(
quantized.reshape(rows, self.kv_lora_rank).view(torch.uint8),
cache_v.to(torch.bfloat16)
.reshape(rows, self.qk_rope_head_dim)
.contiguous()
.view(torch.uint8),
scale.to(torch.float32)
.reshape(rows, self.kv_lora_rank // 128)
.contiguous()
.view(torch.uint8),
),
dim=-1,
)
return packed.view(self.dtype)
def set_kv_buffer(
self,
layer: "RadixAttention",
@@ -709,6 +815,20 @@ class NPUMLATokenToKVPool(MLATokenToKVPool):
loc, _, _ = unwrap_write_loc(loc_info)
self._raise_if_native_kv_cache_disabled()
layer_id = layer.layer_id
if self.dsa_kv_cache_store_fp8:
if cache_v is None:
cache_k, cache_v = cache_k.split(
[self.kv_lora_rank, self.qk_rope_head_dim], dim=-1
)
packed = self._pack_dsa_fp8_kv_cache(cache_k, cache_v)
torch_npu.npu_scatter_nd_update_(
self.k_buffer[layer_id - self.start_layer].view(
-1, 1, self.kv_cache_dim
),
loc.view(-1, 1),
packed.view(-1, 1, self.kv_cache_dim),
)
return
if cache_k.dtype != self.dtype:
cache_k = cache_k.to(self.dtype)
cache_v = cache_v.to(self.dtype)
@@ -748,7 +868,7 @@ class NPUMLATokenToKVPool(MLATokenToKVPool):
index_k = index_k.view(self.store_dtype)
torch_npu.npu_scatter_nd_update_(
self.index_k_buffer[layer_id - self.start_layer].view(
self.index_k_buffer[self._get_indexer_slot(layer_id)].view(
-1, 1, self.index_head_dim
),
loc.view(-1, 1),
@@ -772,19 +892,31 @@ class NPUMLATokenToKVPool(MLATokenToKVPool):
out.append(layer_chunks)
return out
def _get_cpu_offload_layer_buffers(self, local_layer_id):
# flatten(page, slot) also works for the zero-width packed V placeholder.
buffers = [
self.k_buffer[local_layer_id].flatten(0, 1),
self.v_buffer[local_layer_id].flatten(0, 1),
]
slot = self.indexer_layer_id_to_slot.get(local_layer_id + self.start_layer)
if slot is not None:
buffers.append(self.index_k_buffer[slot].flatten(0, 1))
if self.index_k_scale_buffer is not None:
buffers.append(self.index_k_scale_buffer[slot].flatten(0, 1))
if self.dsa_kv_cache_store_fp8:
# Retraction copies opaque records; byte views also avoid FP8
# advanced-indexing restrictions, without decoding/requantizing.
buffers = [
buf.view(torch.uint8) if buf.dtype == torch.float8_e4m3fn else buf
for buf in buffers
]
return buffers
def get_cpu_copy(self, indices, mamba_indices=None, req_pool_index=None):
torch.npu.synchronize()
buf_of_layers = []
has_ik = self.index_head_dim is not None
for local_layer_id in range(self.layer_num):
k_layer = self.k_buffer[local_layer_id].view(-1, 1, self.kv_lora_rank)
v_layer = self.v_buffer[local_layer_id].view(-1, 1, self.qk_rope_head_dim)
ik_layer = (
self.index_k_buffer[local_layer_id].view(-1, 1, self.index_head_dim)
if has_ik
else None
)
buf_of_layers.append([k_layer, v_layer, ik_layer])
buf_of_layers = [
self._get_cpu_offload_layer_buffers(i) for i in range(self.layer_num)
]
kv_cache_cpu = self._chunk_copy_npu_to_cpu(buf_of_layers, indices)
torch.npu.synchronize()
@@ -795,25 +927,12 @@ class NPUMLATokenToKVPool(MLATokenToKVPool):
):
torch.npu.synchronize()
chunk_size = self.cpu_offloading_chunk_size
has_ik = self.index_head_dim is not None
for local_layer_id in range(self.layer_num):
k_layer = self.k_buffer[local_layer_id].view(-1, 1, self.kv_lora_rank)
v_layer = self.v_buffer[local_layer_id].view(-1, 1, self.qk_rope_head_dim)
ik_layer = (
self.index_k_buffer[local_layer_id].view(-1, 1, self.index_head_dim)
if has_ik
else None
)
buffers = self._get_cpu_offload_layer_buffers(local_layer_id)
for i in range(0, len(indices), chunk_size):
chunk_indices = indices[i : i + chunk_size]
chunk = kv_cache_cpu[local_layer_id][i // chunk_size]
k_cpu, v_cpu = chunk[0], chunk[1]
assert k_cpu.shape[0] == len(chunk_indices)
k_layer[chunk_indices] = k_cpu.to(k_layer.device, non_blocking=True)
v_layer[chunk_indices] = v_cpu.to(v_layer.device, non_blocking=True)
if has_ik:
ik_cpu = chunk[2]
ik_layer[chunk_indices] = ik_cpu.to(
ik_layer.device, non_blocking=True
)
for buffer, cpu in zip(buffers, chunk, strict=True):
assert cpu.shape[0] == len(chunk_indices)
buffer[chunk_indices] = cpu.to(buffer.device, non_blocking=True)
torch.npu.synchronize()
@@ -11,6 +11,7 @@ from sglang.srt.hardware_backend.npu.attention.mla_preprocess import (
is_fia_nz,
is_mla_preprocess_enabled,
)
from sglang.srt.hardware_backend.npu.utils import is_npu_arch35
from sglang.srt.layers.attention.dsa.dsa_npu_indexer import scattered_to_tp_attn_full
from sglang.srt.layers.attention.dsa.utils import (
dsa_use_prefill_cp,
@@ -23,6 +24,7 @@ if TYPE_CHECKING:
from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA
from sglang.srt.utils import BumpAllocator
_use_ag_after_qlora = envs.SGLANG_USE_AG_AFTER_QLORA.get()
_is_npu_arch35 = is_npu_arch35()
# region MHA
@@ -347,6 +349,20 @@ def forward_mla_core_npu(
# region DSA
def _apply_interleaved_rope_with_half_output(rotary_emb, positions, q_pe, k_pe):
"""Apply RoPE to interleaved Q/K and return half-layout outputs."""
rotary_emb.get_cos_sin_with_position(positions)
cos = rotary_emb.position_cos.to(device=q_pe.device, dtype=q_pe.dtype).view(
-1, 1, 1, q_pe.shape[-1]
)
sin = rotary_emb.position_sin.to(device=q_pe.device, dtype=q_pe.dtype).view(
-1, 1, 1, q_pe.shape[-1]
)
q_pe = torch_npu.npu_interleave_rope(q_pe.unsqueeze(2), cos, sin).squeeze(2)
k_pe = torch_npu.npu_interleave_rope(k_pe.unsqueeze(2), cos, sin).squeeze(2)
return q_pe, k_pe
def forward_dsa_prepare_npu(
m: "DeepseekV2AttentionMLA",
positions: torch.Tensor,
@@ -357,7 +373,11 @@ def forward_dsa_prepare_npu(
prev_topk_indices: torch.Tensor = None,
):
dynamic_scale = None
if is_mla_preprocess_enabled() and forward_batch.forward_mode.is_decode():
mla_preprocess_used = (
is_mla_preprocess_enabled()
and not forward_batch.forward_mode.is_extend_or_draft_extend_or_mixed()
)
if mla_preprocess_used:
(
q_pe,
k_pe,
@@ -443,16 +463,25 @@ def forward_dsa_prepare_npu(
q_nope, q_pe = q.split([m.qk_nope_head_dim, m.qk_rope_head_dim], dim=-1)
q_nope_out = torch.bmm(q_nope.transpose(0, 1), m.w_kc)
q_nope_out = torch_npu.npu_transpose_batchmatmul(
q_nope,
m.w_kc,
perm_x1=(1, 0, 2),
perm_x2=(0, 1, 2),
perm_y=(1, 0, 2),
)
q_nope_out = q_nope_out.transpose(0, 1)
if m.layer_id == get_token_to_kv_pool().start_layer:
m.rotary_emb.sin_cos_cache = m.rotary_emb.cos_sin_cache.index_select(
0, positions
if is_mla_preprocess_enabled() and not m.rotary_emb.is_neox_style:
# Match the half-layout RoPE outputs used by MLA preprocessing.
q_pe, k_pe = _apply_interleaved_rope_with_half_output(
m.rotary_emb, positions, q_pe, k_pe
)
q_pe, k_pe = m.rotary_emb(positions, q_pe, k_pe)
else:
if m.layer_id == get_token_to_kv_pool().start_layer:
m.rotary_emb.sin_cos_cache = m.rotary_emb.cos_sin_cache.index_select(
0, positions
)
q_pe, k_pe = m.rotary_emb(positions, q_pe, k_pe)
if dsa_use_prefill_cp(forward_batch):
# support allgather+rerrange
@@ -482,6 +511,7 @@ def forward_dsa_prepare_npu(
forward_batch,
zero_allocator,
positions,
mla_preprocess_used,
)
@@ -495,6 +525,7 @@ def forward_dsa_core_npu(
forward_batch: "ForwardBatch",
zero_allocator: "BumpAllocator",
positions: torch.Tensor,
mla_preprocess_used: bool,
# Gated attention (Ling-V3 / BailingMoeV3): the subclass appends its gate
# to inner_state, so every *_core dispatched from forward_core takes it as
# a trailing arg. None everywhere else.
@@ -505,33 +536,31 @@ def forward_dsa_core_npu(
k_nope.contiguous(),
k_nope.contiguous(),
forward_batch,
save_kv_cache=True, # False if forward_batch.forward_mode.is_extend() else True,
save_kv_cache=not mla_preprocess_used,
q_rope=q_pe.contiguous(),
k_rope=k_pe.contiguous(),
topk_indices=topk_indices,
)
attn_output = attn_output.view(-1, m.num_local_heads, m.kv_lora_rank)
attn_bmm_output = torch.empty(
(attn_output.shape[0], m.num_local_heads, m.v_head_dim),
dtype=attn_output.dtype,
device=attn_output.device,
)
if (
if _is_npu_arch35 or (
forward_batch.forward_mode.is_extend()
and not forward_batch.forward_mode.is_draft_extend_v2()
and not forward_batch.forward_mode.is_target_verify()
):
attn_output = attn_output.transpose(0, 1)
torch.bmm(
attn_bmm_output = torch_npu.npu_transpose_batchmatmul(
attn_output,
m.w_vc,
out=attn_bmm_output.view(-1, m.num_local_heads, m.v_head_dim).transpose(
0, 1
),
perm_x1=(1, 0, 2),
perm_x2=(0, 1, 2),
perm_y=(1, 0, 2),
)
else:
attn_bmm_output = torch.empty(
(attn_output.shape[0], m.num_local_heads, m.v_head_dim),
dtype=attn_output.dtype,
device=attn_output.device,
)
attn_output = attn_output.contiguous()
torch.ops.npu.batch_matmul_transpose(attn_output, m.w_vc, attn_bmm_output)
@@ -570,8 +599,11 @@ def npu_mla_preprocess(
m.quant_config,
)
# mlaprolog does not require additional calculation of q_lora
_is_mlaprolog = hasattr(m.quant_config, "ignore") and any(
re.fullmatch(r".*kv_b_proj", l) for l in m.quant_config.ignore
_is_mlaprolog = (
_is_npu_arch35 and get_token_to_kv_pool().index_head_dim is not None
) or (
hasattr(m.quant_config, "ignore")
and any(re.fullmatch(r".*kv_b_proj", l) for l in m.quant_config.ignore)
)
if _is_mlaprolog:
(
@@ -1,5 +1,5 @@
import logging
from typing import TYPE_CHECKING, List, Optional
from typing import TYPE_CHECKING, List, Optional, Tuple
import torch
from torch.nn.parameter import Parameter
@@ -256,22 +256,33 @@ class NPUMXFP8LinearMethod(_NPULinearMethodBase):
def apply(
self,
layer: torch.nn.Module,
x: torch.Tensor,
x: torch.Tensor | Tuple[torch.Tensor, torch.Tensor],
bias: Optional[torch.Tensor] = None,
) -> torch.Tensor:
original_dtype = x.dtype
if original_dtype not in (torch.float16, torch.bfloat16):
x = x.to(torch.bfloat16)
if isinstance(x, tuple):
# MLAProlog supplies a [tokens, hidden] quantized query norm.
qx, input_scale = x
input_shape = qx.shape
if input_scale.dtype == torch.uint8:
input_scale = input_scale.view(_get_float8_e8m0fnu_dtype())
input_scale = input_scale.reshape(
qx.shape[0], qx.shape[1] // (2 * MXFP8_BLOCK_SIZE), 2
).contiguous()
original_dtype = torch.bfloat16
else:
original_dtype = x.dtype
if original_dtype not in (torch.float16, torch.bfloat16):
x = x.to(torch.bfloat16)
original_dtype = torch.bfloat16
# Flatten to 2D [tokens, hidden] for npu_dynamic_mx_quant
input_shape = x.shape
x_2d = x.reshape(-1, x.shape[-1])
# Flatten to 2D [tokens, hidden] for npu_dynamic_mx_quant
input_shape = x.shape
x_2d = x.reshape(-1, x.shape[-1])
# Dynamic MXFP8 activation quantisation
qx, input_scale = torch.ops.npu.npu_dynamic_mx_quant(
x_2d, dst_type=torch.float8_e4m3fn
)
# Dynamic MXFP8 activation quantisation
qx, input_scale = torch.ops.npu.npu_dynamic_mx_quant(
x_2d, dst_type=torch.float8_e4m3fn
)
# MXFP8 matmul (weight & scale already transposed at load time)
# Use the cached FP32 bias from process_weights_after_loading; fall back
@@ -25,7 +25,10 @@ if is_cuda():
deep_gemm = e
if is_npu():
import custom_ops # noqa: F401
try:
import custom_ops # noqa: F401
except ImportError:
pass
from sglang.srt.environ import envs
from sglang.srt.layers import deep_gemm_wrapper
@@ -1,5 +1,7 @@
from __future__ import annotations
from functools import lru_cache
import torch
from sglang.srt.environ import envs
@@ -21,6 +23,36 @@ if is_npu():
_use_ag_after_qlora = envs.SGLANG_USE_AG_AFTER_QLORA.get()
@lru_cache(maxsize=1)
def _create_hadamard_128_cpu() -> torch.Tensor:
matrix = [[1.0]]
while len(matrix) < 128:
matrix = [row + row for row in matrix] + [
row + [-value for value in row] for row in matrix
]
return torch.tensor(matrix, dtype=torch.bfloat16)
def create_npu_hadamard_128(head_dim: int, device) -> torch.Tensor:
assert head_dim == 128
# Match vllm-ascend SFA: BF16 matrix, normalized once on the pool's device.
return (_create_hadamard_128_cpu().to(device=device) / (128**0.5)).contiguous()
def _quantize_npu_indexer_activation(x, hadamard, dst_type):
assert x.dtype == torch.bfloat16 and x.shape[-1] == 128
if x.numel() == 0:
return (
torch.empty_like(x, dtype=dst_type),
torch.empty(x.shape[:-1], dtype=torch.float32, device=x.device),
)
rotated = x @ hadamard
quantized, scale = torch_npu.npu_dynamic_quant(
rotated.reshape(-1, 128), dst_type=dst_type
)
return quantized.reshape(x.shape), scale.to(torch.float32).reshape(x.shape[:-1])
class DSANPUIndexerMixin:
def forward_npu(
self,
@@ -181,9 +213,16 @@ class DSANPUIndexerMixin:
torch.npu.current_stream(),
)
get_token_to_kv_pool().set_index_k_buffer(
layer_id, forward_batch.out_cache_loc, k
)
pool = get_token_to_kv_pool()
use_quant_indexer = pool.index_k_scale_buffer is not None
if use_quant_indexer:
k, k_scale = _quantize_npu_indexer_activation(
k, pool.indexer_hadamard_128, pool.dtype
)
pool.set_index_k_scale_buffer(
layer_id, forward_batch.out_cache_loc, k_scale
)
pool.set_index_k_buffer(layer_id, forward_batch.out_cache_loc, k)
if is_prefill:
if (
self.dsa_enable_prefill_cp
@@ -280,6 +319,32 @@ class DSANPUIndexerMixin:
else block_table
)
if use_quant_indexer:
query, query_scale = _quantize_npu_indexer_activation(
q.view(-1, self.n_heads, self.head_dim),
pool.indexer_hadamard_128,
pool.dtype,
)
topk_indices = torch_npu.npu_quant_lightning_indexer(
query=query,
key=past_key_states,
weights=weights,
query_dequant_scale=query_scale,
key_dequant_scale=pool.get_index_k_scale_buffer(layer_id),
actual_seq_lengths_query=actual_seq_lengths_q.to(torch.int32),
actual_seq_lengths_key=actual_seq_lengths_kv.to(
device=k.device, dtype=torch.int32
),
block_table=block_table,
layout_query="TND",
layout_key="PA_BSND",
sparse_count=self.index_topk,
sparse_mode=3,
query_quant_mode=0,
key_quant_mode=0,
)
return topk_indices.squeeze(1)
topk_indices = torch_npu.npu_lightning_indexer(
query=q.view(-1, self.n_heads, self.head_dim),
key=past_key_states,
@@ -1519,7 +1519,27 @@ class KVCacheConfigurator:
from sglang.srt.hardware_backend.npu.memory_pool_npu import (
NPUMLATokenToKVPool,
)
from sglang.srt.hardware_backend.npu.utils import is_npu_arch35
is_arch35 = is_npu_arch35()
use_compact_indexer_layout = (
is_dsa_model
and is_arch35
and _should_elide_dsa_index_k(is_draft_worker=self.is_draft_worker)
)
indexer_layer_ids = None
if use_compact_indexer_layout:
indexer_layer_ids = tuple(
layer_id
for layer_id in range(
self.layer_info.start_layer,
self.layer_info.end_layer,
)
if not dsa_layer_skips_topk(self.model_config.hf_config, layer_id)
)
use_dsa_fp8_kv_cache_storage = (
self.kv_cache_dtype == torch.float8_e4m3fn and is_arch35
)
token_to_kv_pool = NPUMLATokenToKVPool(
max_total_num_tokens,
page_size=self.pool_page_size,
@@ -1527,6 +1547,14 @@ class KVCacheConfigurator:
kv_lora_rank=self.model_config.kv_lora_rank,
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
index_head_dim=(self.model_config.index_head_dim if is_dsa_model else None),
indexer_layer_ids=indexer_layer_ids,
kv_cache_dim=(
calculate_mla_kv_cache_dim(
model_config=self.model_config, kv_cache_dtype=self.kv_cache_dtype
)
if use_dsa_fp8_kv_cache_storage
else None
),
layer_num=self.layer_info.num_effective_layers,
device=self.device,
enable_memory_saver=get_exec().features.enable_memory_saver,
@@ -55,10 +55,12 @@ from sglang.srt.utils.common import (
ceil_div,
is_float4_e2m1fn_x2,
is_hip,
is_npu,
spec_decode_alloc_len_per_request,
)
_is_hip = is_hip()
_is_npu = is_npu()
@dataclass
@@ -469,6 +471,16 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
element_size = torch._utils._element_size(
DSATokenToKVPool.index_k_with_scale_buffer_dtype
)
if _is_npu:
from sglang.srt.hardware_backend.npu.utils import is_npu_arch35
dtype = kvc.kv_cache_dtype
# GPU sizing above assumes FP8 indexers; NPU also needs BF16 sizing.
if dtype != torch.float8_e4m3fn:
indexer_size_per_token = index_head_dim
element_size = torch._utils._element_size(dtype)
if not is_npu_arch35():
allocate_all_layers = True
memory_config = get_memory()
indexer_ratio = 1
if memory_config.enable_hisparse:
+1 -1
View File
@@ -785,7 +785,7 @@ class DeepseekV2MoE(nn.Module):
not is_packed_weight
and shared_gate_up_weight.dtype == torch.float8_e4m3fn
)
if self.shared_experts_is_fp8:
if self.shared_experts_is_fp8 and not _is_npu:
if (
_use_aiter
and config.quantization_config.get("quant_method")