[XPU][CI] key persistent JIT kernel cache by image content ID (#35337)

This commit is contained in:
ashwini rathi
2026-08-20 10:27:59 +08:00
committed by GitHub
parent 9db4ba8da1
commit 238ba40c27
3 changed files with 45 additions and 11 deletions
@@ -112,6 +112,9 @@ def chunk_gated_delta_rule_fwd_kernel_h_blockdim64_k_loop(
)
index = tl.load(initial_state_indices + i_n).to(tl.int32)
# Padded rows carry the -1 sentinel; without this guard the sentinel
# reaches pointer arithmetic and addresses before the state pool.
valid_state = index >= 0
h0 = initial_state + index * stride_h
ht = initial_state + index * stride_h
if USE_INITIAL_STATE:
@@ -128,18 +131,20 @@ def chunk_gated_delta_rule_fwd_kernel_h_blockdim64_k_loop(
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:
if USE_INITIAL_STATE and valid_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:
elif valid_state:
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)
else:
b_h = tl.zeros([BV, 64], dtype=tl.float32)
# Store pre-update h to output
p_ho = tl.make_block_ptr(
@@ -181,18 +186,20 @@ def chunk_gated_delta_rule_fwd_kernel_h_blockdim64_k_loop(
for k_blk in range(0, K, 64):
# Reload h (same source as Phase 1)
if i_t == 0:
if USE_INITIAL_STATE:
if USE_INITIAL_STATE and valid_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:
elif valid_state:
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)
else:
b_h = tl.zeros([BV, 64], dtype=tl.float32)
# Gate decay on h
if USE_G:
@@ -215,7 +222,7 @@ def chunk_gated_delta_rule_fwd_kernel_h_blockdim64_k_loop(
b_h += tl.trans(tl.dot(b_k, b_v))
# Save updated h to scratch (initial_state) for next time step
if INPLACE_UPDATE:
if INPLACE_UPDATE and valid_state:
p_hs = tl.make_block_ptr(
ht, (V, K), (K, 1), (i_v * BV, k_blk), (BV, 64), (1, 0)
)