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