[AMD] Fuse compress norm+rope+hadamard into single Triton kernel (#26894)

This commit is contained in:
jacky.cheng
2026-06-03 14:20:40 -07:00
committed by GitHub
parent c670609ac5
commit 1dd9432889
2 changed files with 170 additions and 8 deletions
@@ -97,6 +97,7 @@ class CompressorBackendMixin:
from sglang.srt.layers.attention.dsv4.fused_compress_triton import (
hip_compress_forward,
hip_compress_fused_norm_rope_hadamard_inplace,
hip_compress_fused_norm_rope_inplace,
)
@@ -113,14 +114,24 @@ class CompressorBackendMixin:
norm_eps = (
norm.variance_epsilon if hasattr(norm, "variance_epsilon") else norm.eps
)
hip_compress_fused_norm_rope_inplace(
kv_compressed,
norm.weight,
norm_eps,
freqs_cis_cache,
plan,
)
return rotate_activation(kv_compressed) if rotate else kv_compressed
if rotate:
hip_compress_fused_norm_rope_hadamard_inplace(
kv_compressed,
norm.weight,
norm_eps,
freqs_cis_cache,
plan,
head_dim,
)
else:
hip_compress_fused_norm_rope_inplace(
kv_compressed,
norm.weight,
norm_eps,
freqs_cis_cache,
plan,
)
return kv_compressed
kv_compressed = compress_forward(
kv_score_buffer=kv_score_buffer,
@@ -723,6 +723,105 @@ def _compress_norm_rope_kernel(
tl.store(kv_ptr + base + rope_start + 2 * pair_offs + 1, out_imag, mask=pair_mask)
@triton.jit
def _compress_norm_rope_hadamard_kernel(
kv_ptr,
weight_ptr,
freqs_ptr,
handle_ptr,
eps,
hadamard_scale,
kv_row_stride,
freqs_row_stride,
plan_row_stride,
HEAD_DIM: tl.constexpr,
ROPE_DIM: tl.constexpr,
HEAD_BLOCK: tl.constexpr,
ROPE_PAIR_BLOCK: tl.constexpr,
COMPRESS_RATIO: tl.constexpr,
IS_DECODE: tl.constexpr,
LOG2_HEAD_DIM: tl.constexpr,
):
work_id = tl.program_id(0)
if IS_DECODE:
row = work_id
seq_len = tl.load(handle_ptr + work_id).to(tl.int32)
position = ((seq_len - 1) // COMPRESS_RATIO) * COMPRESS_RATIO
else:
plan_base = handle_ptr + work_id * plan_row_stride
row = tl.load(plan_base + 0).to(tl.int32)
plan_position = tl.load(plan_base + 2).to(tl.int32)
if row < 0:
return
position = plan_position + 1 - COMPRESS_RATIO
base = row.to(tl.int64) * kv_row_stride
offs = tl.arange(0, HEAD_BLOCK)
mask = offs < HEAD_DIM
x = tl.load(kv_ptr + base + offs, mask=mask, other=0.0).to(tl.float32)
w = tl.load(weight_ptr + offs, mask=mask, other=0.0).to(tl.float32)
rms_inv = tl.rsqrt(tl.sum(x * x, axis=0) / HEAD_DIM + eps)
x_normed = x * rms_inv * w
rope_start: tl.constexpr = HEAD_DIM - ROPE_DIM
pair_offs = tl.arange(0, ROPE_PAIR_BLOCK)
pair_mask = pair_offs < (ROPE_DIM // 2)
x_real = tl.load(
kv_ptr + base + rope_start + 2 * pair_offs,
mask=pair_mask,
other=0.0,
).to(tl.float32)
x_imag = tl.load(
kv_ptr + base + rope_start + 2 * pair_offs + 1,
mask=pair_mask,
other=0.0,
).to(tl.float32)
w_real = tl.load(
weight_ptr + rope_start + 2 * pair_offs,
mask=pair_mask,
other=1.0,
).to(tl.float32)
w_imag = tl.load(
weight_ptr + rope_start + 2 * pair_offs + 1,
mask=pair_mask,
other=1.0,
).to(tl.float32)
x_real = x_real * rms_inv * w_real
x_imag = x_imag * rms_inv * w_imag
freq_base = position.to(tl.int64) * freqs_row_stride
f_real = tl.load(freqs_ptr + freq_base + 2 * pair_offs, mask=pair_mask, other=0.0)
f_imag = tl.load(
freqs_ptr + freq_base + 2 * pair_offs + 1,
mask=pair_mask,
other=0.0,
)
out_real = x_real * f_real - x_imag * f_imag
out_imag = x_real * f_imag + x_imag * f_real
# Store norm+rope result to kv_ptr (will be used for butterfly stages)
tl.store(kv_ptr + base + offs, x_normed, mask=mask & (offs < rope_start))
tl.store(kv_ptr + base + rope_start + 2 * pair_offs, out_real, mask=pair_mask)
tl.store(kv_ptr + base + rope_start + 2 * pair_offs + 1, out_imag, mask=pair_mask)
# Walsh-Hadamard butterfly transform via store-reload through L1 cache.
# Barriers are required because multiple warps share the same row in memory;
# without them a fast warp can overwrite a partner value before a slow warp reads it.
for stage in tl.static_range(LOG2_HEAD_DIM):
stride = 1 << stage
is_even = ((offs >> stage) & 1) == 0
partner = tl.where(is_even, offs + stride, offs - stride)
tl.debug_barrier()
x_self = tl.load(kv_ptr + base + offs, mask=mask)
x_partner = tl.load(kv_ptr + base + partner, mask=mask)
result = tl.where(is_even, x_self + x_partner, x_partner - x_self)
if stage == LOG2_HEAD_DIM - 1:
result = result * hadamard_scale
tl.debug_barrier()
tl.store(kv_ptr + base + offs, result, mask=mask)
def _plan_as_i32(plan: torch.Tensor) -> torch.Tensor:
assert plan.dtype == torch.uint8 and plan.dim() == 2 and plan.shape[1] == 16
return plan.view(torch.int32).view(-1, 4)
@@ -952,3 +1051,55 @@ def hip_compress_fused_norm_rope_inplace(
COMPRESS_RATIO=plan.compress_ratio,
IS_DECODE=is_decode,
)
def hip_compress_fused_norm_rope_hadamard_inplace(
kv: torch.Tensor,
weight: torch.Tensor,
eps: float,
freqs_cis: torch.Tensor,
plan: Union[CompressorDecodePlan, CompressorPrefillPlan],
head_dim: int,
) -> None:
assert kv.dim() == 2 and kv.stride(-1) == 1
assert weight.shape == (kv.shape[1],)
assert kv.shape[1] == head_dim
freqs_real = torch.view_as_real(freqs_cis).flatten(-2)
rope_dim = freqs_real.shape[-1]
assert head_dim >= rope_dim and rope_dim % 2 == 0
assert (head_dim & (head_dim - 1)) == 0, "head_dim must be power of 2"
is_decode = _is_decode_plan(plan)
if is_decode:
handle = plan.seq_lens
else:
handle = _plan_as_i32(plan.compress_plan)
if handle.numel() == 0:
return
import math
log2_head_dim = int(math.log2(head_dim))
hadamard_scale = head_dim**-0.5
HEAD_BLOCK = triton.next_power_of_2(head_dim)
ROPE_PAIR_BLOCK = max(triton.next_power_of_2(rope_dim // 2), 1)
_compress_norm_rope_hadamard_kernel[(handle.shape[0],)](
kv,
weight,
freqs_real,
handle,
eps,
hadamard_scale,
kv.stride(0),
freqs_real.stride(0),
handle.stride(0) if not is_decode else 0,
HEAD_DIM=head_dim,
ROPE_DIM=rope_dim,
HEAD_BLOCK=HEAD_BLOCK,
ROPE_PAIR_BLOCK=ROPE_PAIR_BLOCK,
COMPRESS_RATIO=plan.compress_ratio,
IS_DECODE=is_decode,
LOG2_HEAD_DIM=log2_head_dim,
)