[NPU] Support GLM-4.7-Flash on NPU (#21408)

This commit is contained in:
Todobe
2026-04-02 17:44:50 +08:00
committed by GitHub
parent 9d9537fbd3
commit 083304ca44
2 changed files with 277 additions and 98 deletions
@@ -21,6 +21,7 @@ from sglang.srt.hardware_backend.npu.attention.mla_preprocess import (
) )
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.layers.attention.nsa.utils import is_nsa_enable_prefill_cp from sglang.srt.layers.attention.nsa.utils import is_nsa_enable_prefill_cp
from sglang.srt.layers.dp_attention import get_attention_tp_size
from sglang.srt.layers.radix_attention import AttentionType from sglang.srt.layers.radix_attention import AttentionType
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.speculative.spec_info import SpecInput from sglang.srt.speculative.spec_info import SpecInput
@@ -217,6 +218,7 @@ class AscendAttnBackend(AttentionBackend):
speculative_step_id + 1, device="npu" speculative_step_id + 1, device="npu"
) )
self.page_size = model_runner.page_size self.page_size = model_runner.page_size
self.model_dtype = model_runner.model_config.dtype
self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA
if self.use_mla: if self.use_mla:
self.kv_lora_rank = model_runner.model_config.kv_lora_rank self.kv_lora_rank = model_runner.model_config.kv_lora_rank
@@ -265,6 +267,18 @@ class AscendAttnBackend(AttentionBackend):
model_runner.token_to_kv_pool.full_to_swa_index_mapping model_runner.token_to_kv_pool.full_to_swa_index_mapping
) )
# head num padding
self.padding_size_list = [1, 2, 4, 8, 16, 32, 64, 128]
self.q_head_num_padding = None
if hasattr(model_runner.model_config, "num_attention_heads") and self.use_mla:
self.tp_q_head_num = (
model_runner.model_config.num_attention_heads // get_attention_tp_size()
)
for num in self.padding_size_list:
if num >= self.tp_q_head_num:
self.q_head_num_padding = num
break
# dllm model config # dllm model config
self.dllm_config = DllmConfig.from_server_args(model_runner.server_args) self.dllm_config = DllmConfig.from_server_args(model_runner.server_args)
self.is_dllm_model = False self.is_dllm_model = False
@@ -449,6 +463,37 @@ class AscendAttnBackend(AttentionBackend):
torch.cumsum(extend_seq_lens_cpu_int, dim=0).int().tolist() torch.cumsum(extend_seq_lens_cpu_int, dim=0).int().tolist()
) )
if (
self.q_head_num_padding is not None
and self.q_head_num_padding > self.tp_q_head_num
):
# In the MLA architecture, the FIA kernel requires the head count to be a power of 2.
# Therefore, we pad the head dimension accordingly and initialize an empty tensor for padding.
metadata.nope_padding = torch.empty(
[
bs,
1,
self.q_head_num_padding - self.tp_q_head_num,
self.kv_lora_rank,
],
dtype=(
self.model_dtype if self.model_dtype is not None else torch.bfloat16
),
device=seq_lens.device,
)
metadata.rope_padding = torch.empty(
[
bs,
1,
self.q_head_num_padding - self.tp_q_head_num,
self.qk_rope_head_dim,
],
dtype=(
self.model_dtype if self.model_dtype is not None else torch.bfloat16
),
device=seq_lens.device,
)
self.graph_metadata[bs] = metadata self.graph_metadata[bs] = metadata
self.forward_metadata = metadata self.forward_metadata = metadata
@@ -1007,110 +1052,212 @@ class AscendAttnBackend(AttentionBackend):
-1, layer.tp_q_head_num * layer.v_head_dim -1, layer.tp_q_head_num * layer.v_head_dim
) )
elif sum(forward_batch.extend_prefix_lens_cpu) > 0: elif sum(forward_batch.extend_prefix_lens_cpu) > 0:
num_token_padding = q.shape[0] # This branch adds support for prefix cache for GLM-4.7-Flash.
q, k, v = [ # When using the MLA architecture, if qk head dim equals v head dim and the head count is not a power of 2,
data[: forward_batch.num_token_non_padded_cpu] for data in [q, k, v] # we use the FIA kernel for computation.
] if layer.qk_head_dim == layer.v_head_dim:
q_nope, q_rope = q.split([layer.v_head_dim, self.qk_rope_head_dim], dim=-1) q = q.reshape(-1, layer.tp_q_head_num, layer.qk_head_dim)
k_nope, k_rope = k.split([layer.v_head_dim, self.qk_rope_head_dim], dim=-1)
# 1st, compute extend tokens to get attn_output and attn_lse k_buffer = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id)
num_tokens = q_nope.size(0) v_buffer = forward_batch.token_to_kv_pool.get_value_buffer(
attn_output = torch.zeros( layer.layer_id
num_tokens, )
layer.tp_q_head_num, kv_cached = torch.index_select(
layer.v_head_dim, k_buffer, 0, self.forward_metadata.flatten_prefix_block_tables
dtype=q_nope.dtype, )
device=q_nope.device, k_rope_cached = torch.index_select(
) v_buffer, 0, self.forward_metadata.flatten_prefix_block_tables
attn_lse = torch.zeros( ).flatten(0, 1)
layer.tp_q_head_num,
num_tokens,
dtype=torch.float32,
device=q_nope.device,
)
torch_npu.atb.npu_ring_mla(
q_nope=q_nope,
q_rope=q_rope,
k_nope=k_nope,
k_rope=k_rope,
value=v,
mask=self.ringmla_mask,
seqlen=self.forward_metadata.extend_seq_lens_cpu_int,
head_num=layer.tp_q_head_num,
kv_head_num=layer.tp_k_head_num,
pre_out=None,
prev_lse=None,
qk_scale=layer.scaling,
kernel_type="kernel_type_high_precision",
mask_type="mask_type_triu",
calc_type="calc_type_first_ring",
output=attn_output,
softmax_lse=attn_lse,
)
# 2nd, load history kvcache(kv_a and k_pe) and calculate k_nope assert layer.kv_b_proj is not None
k_buffer = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id) kv = layer.kv_b_proj(kv_cached)[0].view(
v_buffer = forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id) -1, layer.tp_k_head_num, self.qk_nope_head_dim + layer.v_head_dim
kv_cached = torch.index_select( )
k_buffer, 0, self.forward_metadata.flatten_prefix_block_tables k_nope, v_pre = kv.split(
) [self.qk_nope_head_dim, layer.v_head_dim], dim=-1
k_rope_cached = torch.index_select( )
v_buffer, 0, self.forward_metadata.flatten_prefix_block_tables
).flatten(0, 1)
assert layer.kv_b_proj is not None k_rope = k_rope_cached.expand(-1, layer.tp_k_head_num, -1)
kv = layer.kv_b_proj(kv_cached)[0].view( k_pre = torch.cat([k_nope, k_rope], dim=-1)
-1, layer.tp_k_head_num, self.qk_nope_head_dim + layer.v_head_dim
)
k_nope, v = kv.split([self.qk_nope_head_dim, layer.v_head_dim], dim=-1)
# 3rd, compute history kv to attn_out attn_output = torch.empty(
k_rope = k_rope_cached.expand(-1, layer.tp_k_head_num, -1) (q.size(0), layer.tp_q_head_num, layer.v_head_dim),
seq_len = torch.stack( device=q.device,
[ dtype=q.dtype,
)
q_len_offset = 0
prefix_len_offset = 0
for q_len, prefix_len in zip(
self.forward_metadata.extend_seq_lens_cpu_int, self.forward_metadata.extend_seq_lens_cpu_int,
self.forward_metadata.prefix_lens, self.forward_metadata.prefix_lens,
] ):
) k_cur_slice = k[None, q_len_offset : q_len_offset + q_len]
torch_npu.atb.npu_ring_mla( v_cur_slice = v[None, q_len_offset : q_len_offset + q_len]
q_nope=q_nope, k_pre_slice = k_pre[
q_rope=q_rope, None, prefix_len_offset : prefix_len_offset + prefix_len
k_nope=k_nope, ]
k_rope=k_rope, v_pre_slice = v_pre[
value=v, None, prefix_len_offset : prefix_len_offset + prefix_len
mask=self.ringmla_mask, ]
seqlen=seq_len,
head_num=layer.tp_q_head_num, k_full = torch.cat([k_pre_slice, k_cur_slice], dim=1)
kv_head_num=layer.tp_k_head_num, v_full = torch.cat([v_pre_slice, v_cur_slice], dim=1)
pre_out=attn_output,
prev_lse=attn_lse, attn_output[q_len_offset : q_len_offset + q_len] = (
qk_scale=layer.scaling, torch.ops.npu.npu_fused_infer_attention_score(
kernel_type="kernel_type_high_precision", q[None, q_len_offset : q_len_offset + q_len],
mask_type="no_mask", k_full,
calc_type="calc_type_default", v_full,
output=attn_output, num_heads=layer.tp_q_head_num,
softmax_lse=attn_lse, num_key_value_heads=layer.tp_k_head_num,
) input_layout="BSND", # todo, TND not supports q_heads!=k_heads
attn_output = attn_output.reshape( atten_mask=self.fia_mask,
[-1, layer.tp_q_head_num, layer.v_head_dim] sparse_mode=3,
) scale=layer.scaling,
if num_token_padding != forward_batch.num_token_non_padded_cpu: next_tokens=0,
attn_output = torch.cat( )[0]
[ )
attn_output, q_len_offset += q_len
attn_output.new_zeros( prefix_len_offset += prefix_len
num_token_padding - attn_output.shape[0], attn_output = attn_output.view(
*attn_output.shape[1:], -1, layer.tp_q_head_num * layer.v_head_dim
),
],
dim=0,
) )
else:
num_token_padding = q.shape[0]
q, k, v = [
data[: forward_batch.num_token_non_padded_cpu] for data in [q, k, v]
]
q_nope, q_rope = q.split(
[layer.v_head_dim, self.qk_rope_head_dim], dim=-1
)
k_nope, k_rope = k.split(
[layer.v_head_dim, self.qk_rope_head_dim], dim=-1
)
# 1st, compute extend tokens to get attn_output and attn_lse
num_tokens = q_nope.size(0)
attn_output = torch.zeros(
num_tokens,
layer.tp_q_head_num,
layer.v_head_dim,
dtype=q_nope.dtype,
device=q_nope.device,
)
attn_lse = torch.zeros(
layer.tp_q_head_num,
num_tokens,
dtype=torch.float32,
device=q_nope.device,
)
torch_npu.atb.npu_ring_mla(
q_nope=q_nope,
q_rope=q_rope,
k_nope=k_nope,
k_rope=k_rope,
value=v,
mask=self.ringmla_mask,
seqlen=self.forward_metadata.extend_seq_lens_cpu_int,
head_num=layer.tp_q_head_num,
kv_head_num=layer.tp_k_head_num,
pre_out=None,
prev_lse=None,
qk_scale=layer.scaling,
kernel_type="kernel_type_high_precision",
mask_type="mask_type_triu",
calc_type="calc_type_first_ring",
output=attn_output,
softmax_lse=attn_lse,
)
# 2nd, load history kvcache(kv_a and k_pe) and calculate k_nope
k_buffer = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id)
v_buffer = forward_batch.token_to_kv_pool.get_value_buffer(
layer.layer_id
)
kv_cached = torch.index_select(
k_buffer, 0, self.forward_metadata.flatten_prefix_block_tables
)
k_rope_cached = torch.index_select(
v_buffer, 0, self.forward_metadata.flatten_prefix_block_tables
).flatten(0, 1)
assert layer.kv_b_proj is not None
kv = layer.kv_b_proj(kv_cached)[0].view(
-1, layer.tp_k_head_num, self.qk_nope_head_dim + layer.v_head_dim
)
k_nope, v = kv.split([self.qk_nope_head_dim, layer.v_head_dim], dim=-1)
# 3rd, compute history kv to attn_out
k_rope = k_rope_cached.expand(-1, layer.tp_k_head_num, -1)
seq_len = torch.stack(
[
self.forward_metadata.extend_seq_lens_cpu_int,
self.forward_metadata.prefix_lens,
]
)
torch_npu.atb.npu_ring_mla(
q_nope=q_nope,
q_rope=q_rope,
k_nope=k_nope,
k_rope=k_rope,
value=v,
mask=self.ringmla_mask,
seqlen=seq_len,
head_num=layer.tp_q_head_num,
kv_head_num=layer.tp_k_head_num,
pre_out=attn_output,
prev_lse=attn_lse,
qk_scale=layer.scaling,
kernel_type="kernel_type_high_precision",
mask_type="no_mask",
calc_type="calc_type_default",
output=attn_output,
softmax_lse=attn_lse,
)
attn_output = attn_output.reshape(
[-1, layer.tp_q_head_num, layer.v_head_dim]
)
if num_token_padding != forward_batch.num_token_non_padded_cpu:
attn_output = torch.cat(
[
attn_output,
attn_output.new_zeros(
num_token_padding - attn_output.shape[0],
*attn_output.shape[1:],
),
],
dim=0,
)
else: else:
assert ( if layer.qk_head_dim == layer.v_head_dim:
layer.qk_head_dim != layer.v_head_dim """FIA will support multi-bs in the later version of CANN"""
), "FIA only supports qk_head_dim != v_head_dim" q = q.reshape(-1, layer.tp_q_head_num, layer.qk_head_dim)
if layer.v_head_dim in [256]: attn_output = torch.empty(
(q.size(0), layer.tp_q_head_num, layer.v_head_dim),
device=q.device,
dtype=q.dtype,
)
q_len_offset = 0
for q_len in forward_batch.extend_seq_lens_cpu:
attn_output[q_len_offset : q_len_offset + q_len] = (
torch.ops.npu.npu_fused_infer_attention_score(
q[None, q_len_offset : q_len_offset + q_len],
k[None, q_len_offset : q_len_offset + q_len],
v[None, q_len_offset : q_len_offset + q_len],
num_heads=layer.tp_q_head_num,
num_key_value_heads=layer.tp_k_head_num,
input_layout="BSND", # todo, TND not supports q_heads!=k_heads
atten_mask=self.fia_mask.unsqueeze(0),
sparse_mode=3 if q_len != 1 else 0,
scale=layer.scaling,
next_tokens=0,
)[0]
)
q_len_offset += q_len
attn_output = attn_output.view(
-1, layer.tp_q_head_num * layer.v_head_dim
)
elif layer.v_head_dim in [256]:
"""Currently, in NO_QUANT situation, qk_nope_head_dim == v_head_dim, and rope exists, v_head_dim only support 512 and 128""" """Currently, in NO_QUANT situation, qk_nope_head_dim == v_head_dim, and rope exists, v_head_dim only support 512 and 128"""
kv_lora_rank = k.shape[-1] - self.qk_rope_head_dim kv_lora_rank = k.shape[-1] - self.qk_rope_head_dim
kv_c, k_rope = k.split([kv_lora_rank, self.qk_rope_head_dim], dim=-1) kv_c, k_rope = k.split([kv_lora_rank, self.qk_rope_head_dim], dim=-1)
@@ -1543,6 +1690,24 @@ class AscendAttnBackend(AttentionBackend):
q_nope = q.view(-1, 1, layer.tp_q_head_num, self.kv_lora_rank).contiguous() q_nope = q.view(-1, 1, layer.tp_q_head_num, self.kv_lora_rank).contiguous()
q_rope = q_rope.view(-1, 1, layer.tp_q_head_num, self.qk_rope_head_dim) q_rope = q_rope.view(-1, 1, layer.tp_q_head_num, self.qk_rope_head_dim)
assert (
self.q_head_num_padding is None
or self.q_head_num_padding >= layer.tp_q_head_num
)
if (
self.q_head_num_padding is not None
and self.q_head_num_padding > layer.tp_q_head_num
):
# The FIA kernel only supports head counts that are powers of 2.
# Therefore, we pad the head dimension when it is not a power of 2.
q_nope = torch.cat(
[q_nope, self.forward_metadata.nope_padding], dim=2
).contiguous()
q_rope = torch.cat(
[q_rope, self.forward_metadata.rope_padding], dim=2
).contiguous()
if self.forward_metadata.seq_lens_cpu_int is None: if self.forward_metadata.seq_lens_cpu_int is None:
actual_seq_len_kv = self.forward_metadata.seq_lens_cpu_list actual_seq_len_kv = self.forward_metadata.seq_lens_cpu_list
else: else:
@@ -1556,7 +1721,7 @@ class AscendAttnBackend(AttentionBackend):
c_kv_cache, c_kv_cache,
query_rope=q_rope, query_rope=q_rope,
key_rope=k_rope_cache, key_rope=k_rope_cache,
num_heads=layer.tp_q_head_num, num_heads=self.q_head_num_padding,
num_key_value_heads=layer.tp_k_head_num, num_key_value_heads=layer.tp_k_head_num,
block_table=self.forward_metadata.block_tables, block_table=self.forward_metadata.block_tables,
block_size=self.page_size, block_size=self.page_size,
@@ -1576,7 +1741,7 @@ class AscendAttnBackend(AttentionBackend):
c_kv_cache, c_kv_cache,
query_rope=q_rope, query_rope=q_rope,
key_rope=k_rope_cache, key_rope=k_rope_cache,
num_heads=layer.tp_q_head_num, num_heads=self.q_head_num_padding,
num_key_value_heads=layer.tp_k_head_num, num_key_value_heads=layer.tp_k_head_num,
block_table=self.forward_metadata.block_tables, block_table=self.forward_metadata.block_tables,
block_size=self.page_size, block_size=self.page_size,
@@ -1589,6 +1754,8 @@ class AscendAttnBackend(AttentionBackend):
workspace=workspace, workspace=workspace,
out=[output, softmax_lse], out=[output, softmax_lse],
) )
output = output[:, :, : layer.tp_q_head_num, :]
return output.view(-1, layer.tp_q_head_num * self.kv_lora_rank) return output.view(-1, layer.tp_q_head_num * self.kv_lora_rank)
def forward_decode( def forward_decode(
@@ -271,7 +271,16 @@ class RotaryEmbedding(MultiPlatformOp):
rotary_mode = "half" rotary_mode = "half"
else: else:
rotary_mode = "interleave" rotary_mode = "interleave"
mrope_section = [0, 0, 0] mrope_section = [0, 0, 0]
# The npu_mrope kernel only supports 1D or 2D tensors for query and key.
# Therefore, when their dimensions exceed 2D, we flatten query and key to 2D tensors before computation
# and reshape their original shapes afterward.
query_shape = query.shape
key_shape = key.shape
query = query.reshape(query.shape[0], -1)
key = key.reshape(key.shape[0], -1)
query_out, key_out = torch_npu.npu_mrope( query_out, key_out = torch_npu.npu_mrope(
positions, positions,
query, query,
@@ -281,6 +290,9 @@ class RotaryEmbedding(MultiPlatformOp):
mrope_section=mrope_section, mrope_section=mrope_section,
rotary_mode=rotary_mode, rotary_mode=rotary_mode,
) )
query_out = query_out.reshape(query_shape)
key_out = key_out.reshape(key_shape)
return query_out, key_out return query_out, key_out
def forward_cpu( def forward_cpu(