[AMD] Pad QSA MQA decode Q-heads to 16 for ROCm MFMA (#38875)
This commit is contained in:
@@ -6,6 +6,8 @@ from typing import Optional
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.utils.common import is_hip
|
||||||
|
|
||||||
try:
|
try:
|
||||||
import flashinfer.comm # noqa: F401
|
import flashinfer.comm # noqa: F401
|
||||||
except ImportError:
|
except ImportError:
|
||||||
@@ -349,10 +351,14 @@ def tilelang_qsa_mqa_decode(
|
|||||||
)
|
)
|
||||||
if not q.shape[0] or not max_model_len:
|
if not q.shape[0] or not max_model_len:
|
||||||
return logits
|
return logits
|
||||||
# The validated MMA layout requires N (the Q-head dimension) to be a
|
# CUDA MMA accepts an eight-wide N dimension, while ROCm MFMA requires
|
||||||
# multiple of eight. Zero-padding preserves the weight-free head sum.
|
# sixteen. Zero-padding preserves the weight-free head sum on both paths.
|
||||||
query_heads, head_dim = q.shape[1:]
|
query_heads, head_dim = q.shape[1:]
|
||||||
kernel_heads = max(8, ((query_heads + 7) // 8) * 8)
|
head_alignment = 16 if is_hip() else 8
|
||||||
|
kernel_heads = max(
|
||||||
|
head_alignment,
|
||||||
|
((query_heads + head_alignment - 1) // head_alignment) * head_alignment,
|
||||||
|
)
|
||||||
q_kernel = q.to(torch.bfloat16)
|
q_kernel = q.to(torch.bfloat16)
|
||||||
if kernel_heads != query_heads:
|
if kernel_heads != query_heads:
|
||||||
q_kernel = torch.cat(
|
q_kernel = torch.cat(
|
||||||
|
|||||||
Reference in New Issue
Block a user