ironic-mirror / python /xrex_unified /triton_kernels.py
SNAPKITTYWEST's picture
push from SNAPKITTYWEST/ironic-mirror
677e207 verified
Raw
History Blame Contribute Delete
10.1 kB
#
# Copyright (c) 2026 BEL ESPRIT D ACCORD TRUST HOLDINGS INC
# All rights reserved.
# SPDX-License-Identifier: Apache-2.0
# Copyright 2026 X.AI Corp.
"""
Pallas/Triton forward and backward kernels for unified attention.
"""
import jax
import jax.numpy as jnp
from jax import lax
from jax.experimental import pallas as pl
from jax.experimental.pallas import triton as pltriton
from .cap_functions import cap_forward, cap_grad
from .kernel_config import KernelConfig
from .softmax_state import SoftmaxState
def make_triton_forward_kernel(config: KernelConfig):
def kernel(q_ref, k_ref, v_ref, temp_ref, segment_ref,
o_ref, *residual_refs):
seq_len = q_ref.shape[0]
head_dim = q_ref.shape[1]
block_q = config.block_q
block_kv = config.block_kv
start_q = pl.program_id(0)
q = q_ref[pl.dslice(start_q * block_q, block_q), :]
seg_q = segment_ref[pl.dslice(start_q * block_q, block_q)]
seg_q = jnp.expand_dims(seg_q, axis=-1)
state = SoftmaxState.init(block_q, head_dim)
def body(start_k, carry):
state = carry
k = k_ref[:, pl.dslice(start_k * block_kv, block_kv)]
seg_k = segment_ref[pl.dslice(start_k * block_kv, block_kv)]
v = v_ref[pl.dslice(start_k * block_kv, block_kv), :]
temp = temp_ref[pl.dslice(start_q * block_q, block_q)]
temp = jnp.expand_dims(temp, axis=-1)
qk = pl.dot(q, k)
if config.sm_scale != 1.0:
qk *= config.sm_scale
mask = jnp.equal(seg_q, jnp.expand_dims(seg_k, axis=-2))
if config.causal or config.window_len > 0:
span_q = start_q * block_q + jnp.arange(block_q)
span_k = start_k * block_kv + jnp.arange(block_kv)
if config.causal:
causal_mask = span_q[:, None] >= span_k[None, :]
mask = mask & causal_mask
if config.window_len > 0:
window_mask = span_k[None, :] > span_q[:, None] - config.window_len
mask = mask & window_mask
qk = jnp.where(mask, qk, -jnp.inf)
state = state.update(qk, v, config.cap_method, config.cap_params, temp)
return state
if config.causal:
upper_bound = lax.div(start_q * block_q, block_kv) + 1
else:
upper_bound = pl.cdiv(seq_len, block_kv)
if config.window_len > 0:
max_blocks = lax.div(config.window_len + block_q, block_kv)
lower_bound = jnp.maximum(upper_bound - max_blocks, 0)
else:
lower_bound = 0
state = lax.fori_loop(lower_bound, upper_bound, body, state)
if residual_refs:
l_ref, m_ref = residual_refs
l_ref[pl.ds(start_q * block_q, block_q)] = state.l
m_ref[pl.ds(start_q * block_q, block_q)] = state.m
out = state.finalize()
o_ref[pl.dslice(start_q * block_q, block_q), :] = out.astype(o_ref.dtype)
return kernel
def make_triton_backward_kernel_dq(config: KernelConfig):
def kernel(q_ref, k_ref, v_ref, temp_ref, segment_ref,
out_ref, do_scaled_ref, l_ref, m_ref, delta_ref, dq_ref):
seq_len = q_ref.shape[0]
head_dim = q_ref.shape[1]
block_q = config.block_q
block_kv = config.block_kv
start_q = pl.program_id(2)
q = pl.load(q_ref, (pl.ds(start_q * block_q, block_q), slice(None)))
span_q = start_q * block_q + jnp.arange(block_q)
m = pl.load(m_ref, (pl.ds(start_q * block_q, block_q),))
l = pl.load(l_ref, (pl.ds(start_q * block_q, block_q),))
do = pl.load(do_scaled_ref, (pl.ds(start_q * block_q, block_q), slice(None)))
di = pl.load(delta_ref, (pl.ds(start_q * block_q, block_q),))
dq = jnp.zeros([block_q, head_dim], dtype=jnp.float32)
seg_q = pl.load(segment_ref, (pl.ds(start_q * block_q, block_q),))
seg_q = jnp.expand_dims(seg_q, axis=-1)
def inner_loop(start_k, dq):
k = pl.load(k_ref, (pl.ds(start_k * block_kv, block_kv), slice(None)))
v = pl.load(v_ref, (pl.ds(start_k * block_kv, block_kv), slice(None)))
seg_k = pl.load(segment_ref, (pl.dslice(start_k * block_kv, block_kv),))
mask = jnp.equal(seg_q, jnp.expand_dims(seg_k, axis=-2))
temp = pl.load(temp_ref, (pl.dslice(start_q * block_q, block_q),))
temp = jnp.expand_dims(temp, axis=-1)
qk = jnp.zeros((block_q, block_kv), dtype=jnp.float32)
qk += pl.dot(q, k.T)
if config.sm_scale != 1.0:
qk *= config.sm_scale
qk_capped = cap_forward(qk, config.cap_method, config.cap_params)
qk_capped *= temp
span_k = start_k * block_kv + jnp.arange(block_kv)
if config.causal:
causal_mask = span_q[:, None] >= span_k[None, :]
mask = mask & causal_mask
if config.window_len > 0:
window_mask = span_k[None, :] > span_q[:, None] - config.window_len
mask = mask & window_mask
qk_capped = jnp.where(mask, qk_capped, -jnp.inf)
p = jnp.exp(qk_capped - m[:, None])
p = p / jnp.sum(p, axis=1, keepdims=True)
dp = -di[:, None] + pl.dot(do, v.T)
ds = p * dp
if config.z_loss_weight > 0:
log_l = jnp.log(l + 1e-12)
ds += config.z_loss_weight * p * ((log_l + m) / l)[:, None]
ds *= temp
cap_g = cap_grad(qk, config.cap_method, config.cap_params, qk_capped)
ds *= cap_g
if config.sm_scale != 1.0:
ds *= config.sm_scale
dq += pl.dot(ds.astype(k.dtype), k).astype(dq.dtype)
return dq
if config.causal:
upper_bound = lax.div(start_q * block_q, block_kv) + 1
else:
upper_bound = pl.cdiv(seq_len, block_kv)
if config.window_len > 0:
max_blocks = lax.div(config.window_len + block_q, block_kv)
lower_bound = jnp.maximum(upper_bound - max_blocks, 0)
else:
lower_bound = 0
dq = lax.fori_loop(lower_bound, upper_bound, inner_loop, dq)
pl.store(dq_ref, (pl.ds(start_q * block_q, block_q), slice(None)), dq)
return kernel
def make_triton_backward_kernel_dkv(config: KernelConfig):
def kernel(q_ref, k_ref, v_ref, temp_ref, segment_ref,
out_ref, do_scaled_ref, l_ref, m_ref, delta_ref,
dk_ref, dv_ref):
seq_len = q_ref.shape[0]
head_dim = q_ref.shape[1]
block_q = config.block_q
block_kv = config.block_kv
start_k = pl.program_id(2)
dv = jnp.zeros([block_kv, head_dim], dtype=jnp.float32)
dk = jnp.zeros([block_kv, head_dim], dtype=jnp.float32)
k = pl.load(k_ref, (pl.ds(start_k * block_kv, block_kv), slice(None)))
v = pl.load(v_ref, (pl.ds(start_k * block_kv, block_kv), slice(None)))
span_k = start_k * block_kv + jnp.arange(block_kv)
seg_k = pl.load(segment_ref, (pl.ds(start_k * block_kv, block_kv),))
seg_k = jnp.expand_dims(seg_k, axis=-2)
def inner_loop(start_q, carry):
dv, dk = carry
q = pl.load(q_ref, (pl.ds(start_q * block_q, block_q), slice(None)))
qk = jnp.zeros((block_q, block_kv), dtype=jnp.float32)
qk += pl.dot(q, k.T)
seg_q = pl.load(segment_ref, (pl.ds(start_q * block_q, block_q),))
mask = jnp.equal(jnp.expand_dims(seg_q, axis=-1), seg_k)
temp = pl.load(temp_ref, (pl.dslice(start_q * block_q, block_q),))
temp = jnp.expand_dims(temp, axis=-1)
if config.sm_scale != 1.0:
qk *= config.sm_scale
qk_capped = cap_forward(qk, config.cap_method, config.cap_params)
qk_capped *= temp
span_q = start_q * block_q + jnp.arange(block_q)
if config.causal:
causal_mask = span_q[:, None] >= span_k[None, :]
mask = mask & causal_mask
if config.window_len > 0:
window_mask = span_k[None, :] > span_q[:, None] - config.window_len
mask = mask & window_mask
qk_capped = jnp.where(mask, qk_capped, -jnp.inf)
m = pl.load(m_ref, (pl.ds(start_q * block_q, block_q),))
p = jnp.exp(qk_capped - m[:, None])
p = p / jnp.sum(p, axis=1, keepdims=True)
do = pl.load(do_scaled_ref, (pl.ds(start_q * block_q, block_q), slice(None)))
dv += pl.dot(p.astype(do.dtype).T, do)
di = pl.load(delta_ref, (pl.ds(start_q * block_q, block_q),))
dp = -di[:, None] + pl.dot(do, v.T)
ds = p * dp
if config.z_loss_weight > 0:
l = pl.load(l_ref, (pl.ds(start_q * block_q, block_q),))
log_l = jnp.log(l + 1e-12)
ds += config.z_loss_weight * p * ((log_l + m) / l)[:, None]
ds *= temp
cap_g = cap_grad(qk, config.cap_method, config.cap_params, qk_capped)
ds *= cap_g
if config.sm_scale != 1.0:
ds *= config.sm_scale
dk += pl.dot(ds.astype(q_ref.dtype).T, q)
return dv, dk
if config.causal:
lower_bound = lax.div(start_k * block_kv, block_q)
else:
lower_bound = 0
if config.window_len > 0:
max_blocks = lax.div(config.window_len + block_kv, block_q)
upper_bound = jnp.minimum(lower_bound + max_blocks, pl.cdiv(seq_len, block_q))
else:
upper_bound = pl.cdiv(seq_len, block_q)
dv, dk = lax.fori_loop(lower_bound, upper_bound, inner_loop, (dv, dk))
pl.store(dv_ref, (pl.ds(start_k * block_kv, block_kv), slice(None)), dv.astype(dv_ref.dtype))
pl.store(dk_ref, (pl.ds(start_k * block_kv, block_kv), slice(None)), dk.astype(dk_ref.dtype))
return kernel