[XPU] fix correctness issue of GDN triton kernel for XPU (#26065)

Co-authored-by: Ma Mingfei <mingfei.ma@intel.com>
This commit is contained in:
Xia Weiwen
2026-05-25 13:18:11 +08:00
committed by GitHub
co-authored by Ma Mingfei
parent ec6fcb93cb
commit 2bd3ac0b5d
3 changed files with 102 additions and 42 deletions
@@ -16,7 +16,10 @@ from sglang.srt.layers.attention.fla.utils import (
CHUNK_SIZE = 64
# This kernel handles K blocks in a for loop to minimize register spills
# This kernel handles K blocks in a for loop to minimize register spills.
# Time is the OUTER loop; K blocks are processed in two inner phases per step:
# Phase 1: store h to output, accumulate v_correction = sum_k(w_k @ h_k^T)
# Phase 2: update h = gate * h + k^T @ v_gated, save to scratch (initial_state)
@triton.autotune(
configs=[triton.Config({"BV": 64}, num_warps=8, num_stages=2)],
key=["H", "K", "V", "BT", "USE_GK", "USE_INITIAL_STATE", "NT_BUCKET"],
@@ -110,67 +113,104 @@ def chunk_gated_delta_rule_fwd_kernel_h_blockdim64_k_loop(
if INPLACE_UPDATE:
ht = ht + i_h * V * K
# Explicit K loop here to reduce register pressure
for k_start in range(0, K, 64):
# [BV, BK]
b_h1 = tl.zeros([BV, 64], dtype=tl.float32)
# main recurrence — time is the outer loop
for i_t in range(NT):
########################################################################
# Phase 1: store h to output, compute v_new = u - sum_k(w_k @ h_k^T)
########################################################################
b_v_corr = tl.zeros([BT, BV], dtype=tl.float32)
for k_blk in range(0, K, 64):
# Load h: from initial_state (i_t==0) or scratch (i_t>0)
if i_t == 0:
if USE_INITIAL_STATE:
p_hs = tl.make_block_ptr(
h0, (V, K), (K, 1), (i_v * BV, k_blk), (BV, 64), (1, 0)
)
b_h = tl.load(p_hs, boundary_check=(0, 1)).to(tl.float32)
else:
b_h = tl.zeros([BV, 64], dtype=tl.float32)
else:
p_hs = tl.make_block_ptr(
ht, (V, K), (K, 1), (i_v * BV, k_blk), (BV, 64), (1, 0)
)
b_h = tl.load(p_hs, boundary_check=(0, 1)).to(tl.float32)
# load initial state
if USE_INITIAL_STATE:
p_h0_1 = tl.make_block_ptr(
h0, (V, K), (K, 1), (i_v * BV, k_start), (BV, 64), (1, 0)
)
b_h1 += tl.load(p_h0_1, boundary_check=(0, 1)).to(tl.float32)
# main recurrence
for i_t in range(NT):
p_h1 = tl.make_block_ptr(
# Store pre-update h to output
p_ho = tl.make_block_ptr(
h + i_t * stride_h,
(V, K),
(K, 1),
(i_v * BV, k_start),
(i_v * BV, k_blk),
(BV, 64),
(1, 0),
)
tl.store(p_h1, b_h1.to(p_h1.dtype.element_ty), boundary_check=(0, 1))
tl.store(p_ho, b_h.to(p_ho.dtype.element_ty), boundary_check=(0, 1))
b_w = w_desc.load([i_t * BT, k_start])
b_v = tl.dot(b_w, tl.trans(b_h1).to(b_w.dtype))
b_v = v_desc.load([i_t * BT, i_v * BV]) - b_v
# Accumulate correction: w_k @ h_k^T
b_w = w_desc.load([i_t * BT, k_blk])
b_v_corr += tl.dot(b_w, tl.trans(b_h).to(b_w.dtype))
if SAVE_NEW_VALUE:
v_new_desc.store([i_t * BT, i_v * BV], b_v.to(v_new.dtype.element_ty))
# v_new = u - correction
b_v = v_desc.load([i_t * BT, i_v * BV]) - b_v_corr
last_idx = min((i_t + 1) * BT, T) - 1
if USE_G:
b_g_last = tl.load(g + bos * H + last_idx * H + i_h)
p_g = tl.make_block_ptr(
g + bos * H + i_h, (T,), (H,), (i_t * BT,), (BT,), (0,)
if SAVE_NEW_VALUE:
v_new_desc.store([i_t * BT, i_v * BV], b_v.to(v_new.dtype.element_ty))
# Apply gate to v
last_idx = min((i_t + 1) * BT, T) - 1
if USE_G:
b_g_last = tl.load(g + bos * H + last_idx * H + i_h)
p_g = tl.make_block_ptr(
g + bos * H + i_h, (T,), (H,), (i_t * BT,), (BT,), (0,)
)
b_g = tl.load(p_g, boundary_check=(0,))
b_v = b_v * tl.expand_dims(safe_exp(b_g_last - b_g), 1)
b_g_last = exp(b_g_last)
b_v = b_v.to(k.dtype.element_ty)
########################################################################
# Phase 2: reload h, apply gate, update h += k^T @ v, save to scratch
########################################################################
for k_blk in range(0, K, 64):
# Reload h (same source as Phase 1)
if i_t == 0:
if USE_INITIAL_STATE:
p_hs = tl.make_block_ptr(
h0, (V, K), (K, 1), (i_v * BV, k_blk), (BV, 64), (1, 0)
)
b_h = tl.load(p_hs, boundary_check=(0, 1)).to(tl.float32)
else:
b_h = tl.zeros([BV, 64], dtype=tl.float32)
else:
p_hs = tl.make_block_ptr(
ht, (V, K), (K, 1), (i_v * BV, k_blk), (BV, 64), (1, 0)
)
b_g = tl.load(p_g, boundary_check=(0,))
b_v = b_v * safe_exp(b_g_last - b_g)[:, None]
b_g_last = exp(b_g_last)
b_h1 = b_h1 * b_g_last
b_h = tl.load(p_hs, boundary_check=(0, 1)).to(tl.float32)
# Gate decay on h
if USE_G:
b_h = b_h * b_g_last
if USE_GK:
o_k1 = tl.arange(0, 64) + k_start
o_k1 = tl.arange(0, 64) + k_blk
b_gk_last1 = tl.load(
gk + (bos + last_idx) * H * K + i_h * K + o_k1,
mask=(o_k1 < K),
other=0.0,
)
b_h1 *= exp(b_gk_last1)[None, :]
b_v = b_v.to(k.dtype.element_ty)
b_h *= tl.expand_dims(exp(b_gk_last1), 0)
b_k = tl.trans(k_desc.load([i_t * BT, k_start]))
b_h1 += tl.trans(tl.dot(b_k, b_v))
# Delta update: h += k^T @ v
b_k = tl.trans(k_desc.load([i_t * BT, k_blk]))
b_h += tl.trans(tl.dot(b_k, b_v))
# epilogue
if INPLACE_UPDATE:
p_ht = tl.make_block_ptr(
ht, (V, K), (K, 1), (i_v * BV, k_start), (BV, 64), (1, 0)
)
tl.store(p_ht, b_h1.to(p_ht.dtype.element_ty), boundary_check=(0, 1))
# Save updated h to scratch (initial_state) for next time step
if INPLACE_UPDATE:
p_hs = tl.make_block_ptr(
ht, (V, K), (K, 1), (i_v * BV, k_blk), (BV, 64), (1, 0)
)
tl.store(p_hs, b_h.to(p_hs.dtype.element_ty), boundary_check=(0, 1))
def chunk_gated_delta_rule_fwd_h(
@@ -233,6 +233,16 @@ def chunk_gated_delta_rule_fwd_kkt_solve_kernel_low_reg(
)
tl.store(p_Ai_ij, b_Ai_ij.to(A.dtype.element_ty), boundary_check=(0, 1))
# Clean up scratch slots: Pass 2 stored raw A_ij blocks in the upper-triangular
# part of row i_tc0 (cols BC..3*BC). These must be zeroed because
# recompute_w_u_fwd reads the full BT×BT block.
b_zero = tl.zeros([BC, BC], dtype=tl.float32)
for sc in tl.static_range(1, BT // BC):
p_scratch = tl.make_block_ptr(
A, (T, BT), (H * BT, 1), (i_tc0, sc * BC), (BC, BC), (1, 0)
)
tl.store(p_scratch, b_zero.to(A.dtype.element_ty), boundary_check=(0, 1))
def chunk_gated_delta_rule_fwd_intra(
k: torch.Tensor,