[Not-Merge][AMD] GLM-5 performance optimization (#21166)

This commit is contained in:
wufann
2026-04-11 23:58:11 -07:00
committed by GitHub
parent edaa5973d4
commit 19cb918653
2 changed files with 48 additions and 8 deletions
@@ -496,9 +496,6 @@ class Indexer(MultiPlatformOp):
max_seq_len, max_seq_len,
Preshuffle=False, Preshuffle=False,
KVBlockSize=block_kv, KVBlockSize=block_kv,
ChunkK=128,
TotalCuCount=256,
WavePerEU=5,
) )
else: else:
logits = deep_gemm.fp8_paged_mqa_logits( logits = deep_gemm.fp8_paged_mqa_logits(
@@ -342,6 +342,12 @@ class NativeSparseAttnBackend(
dtype=torch.int32, dtype=torch.int32,
device=self.device, device=self.device,
) )
# Aiter mla_decode_fwd supports num_heads multiples of 16 in range [16, 128].
# For models with fewer heads per GPU (e.g. GLM-5 64 heads / TP8 = 8), need to pad the heads to 16.
self.need_pad_heads = self.num_q_heads < 16
self.head_repeat_factor = (
16 // self.num_q_heads if self.num_q_heads < 16 else 1
)
# Speculative decoding # Speculative decoding
self.topk = model_runner.server_args.speculative_eagle_topk or 0 self.topk = model_runner.server_args.speculative_eagle_topk or 0
@@ -1844,6 +1850,21 @@ class NativeSparseAttnBackend(
else: else:
o = torch.empty_like(q) o = torch.empty_like(q)
if self.need_pad_heads:
q_kernel = q.view(
-1, layer.tp_q_head_num, layer.head_dim
).repeat_interleave(self.head_repeat_factor, dim=1)
o_kernel = q.new_empty(
(
q.shape[0],
layer.tp_q_head_num * self.head_repeat_factor,
layer.v_head_dim,
)
)
else:
q_kernel = q.view(-1, layer.tp_q_head_num, layer.head_dim)
o_kernel = o.view(-1, layer.tp_q_head_num, layer.v_head_dim)
kv_indptr = self.kv_indptr kv_indptr = self.kv_indptr
non_minus1_mask = page_table_1 != -1 non_minus1_mask = page_table_1 != -1
@@ -1854,9 +1875,9 @@ class NativeSparseAttnBackend(
get_valid_kv_indices(page_table_1, kv_indptr, kv_indices, bs) get_valid_kv_indices(page_table_1, kv_indptr, kv_indices, bs)
mla_decode_fwd( mla_decode_fwd(
q.view(-1, layer.tp_q_head_num, layer.head_dim), q_kernel,
kv_cache.view(-1, 1, 1, layer.head_dim), kv_cache.view(-1, 1, 1, layer.head_dim),
o.view(-1, layer.tp_q_head_num, layer.v_head_dim), o_kernel,
metadata.cu_seqlens_q, metadata.cu_seqlens_q,
kv_indptr, kv_indptr,
kv_indices, kv_indices,
@@ -1865,7 +1886,10 @@ class NativeSparseAttnBackend(
sm_scale=layer.scaling, sm_scale=layer.scaling,
logit_cap=layer.logit_cap, logit_cap=layer.logit_cap,
) )
# kv_cache = kv_cache.view(-1, 1, layer.head_dim)
if self.need_pad_heads:
o = o_kernel[:, :: self.head_repeat_factor, :]
return o return o
def _forward_aiter_extend( def _forward_aiter_extend(
@@ -1883,6 +1907,21 @@ class NativeSparseAttnBackend(
else: else:
o = torch.empty_like(q) o = torch.empty_like(q)
if self.need_pad_heads:
q_kernel = q.view(
-1, layer.tp_q_head_num, layer.head_dim
).repeat_interleave(self.head_repeat_factor, dim=1)
o_kernel = q.new_empty(
(
num_tokens,
layer.tp_q_head_num * self.head_repeat_factor,
layer.v_head_dim,
)
)
else:
q_kernel = q.view(-1, layer.tp_q_head_num, layer.head_dim)
o_kernel = o.view(-1, layer.tp_q_head_num, layer.v_head_dim)
non_minus1_mask = page_table_1 != -1 non_minus1_mask = page_table_1 != -1
non_minus1_counts = non_minus1_mask.sum(dim=1) non_minus1_counts = non_minus1_mask.sum(dim=1)
@@ -1904,9 +1943,9 @@ class NativeSparseAttnBackend(
) )
# TODO support more forward_mode # TODO support more forward_mode
mla_decode_fwd( mla_decode_fwd(
q.view(-1, layer.tp_q_head_num, layer.head_dim), q_kernel,
kv_cache.view(-1, 1, 1, layer.head_dim), kv_cache.view(-1, 1, 1, layer.head_dim),
o.view(-1, layer.tp_q_head_num, layer.v_head_dim), o_kernel,
cu_seqlens_q, cu_seqlens_q,
kv_indptr, kv_indptr,
kv_indices, kv_indices,
@@ -1915,6 +1954,10 @@ class NativeSparseAttnBackend(
sm_scale=layer.scaling, sm_scale=layer.scaling,
logit_cap=layer.logit_cap, logit_cap=layer.logit_cap,
) )
if self.need_pad_heads:
o = o_kernel[:, :: self.head_repeat_factor, :]
return o return o
def _forward_trtllm( def _forward_trtllm(