# # 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