| |
| |
| |
|
|
| |
| |
| """ |
| Public API: Unified Attention Function. |
| |
| Supports both Triton (Pallas) and Mosaic GPU backends with: |
| - All 4 cap methods (tanh, soft_sign, alsc, none) |
| - Ranker segment bounds + causal/window composition |
| - Online softmax with exact backward |
| - GQA (grouped query attention) |
| """ |
|
|
| import functools |
|
|
| import jax |
| import jax.numpy as jnp |
| from jax.experimental import pallas as pl |
| from jax.experimental.pallas import triton as pltriton |
|
|
| from .cap_functions import cap_forward, cap_grad |
| from .segment_bounds import SegmentBounds |
| from .kernel_config import KernelConfig |
| from .triton_kernels import ( |
| make_triton_forward_kernel, |
| make_triton_backward_kernel_dq, |
| make_triton_backward_kernel_dkv, |
| ) |
|
|
|
|
| @functools.partial(jax.custom_vjp, nondiff_argnums=(5,)) |
| @functools.partial(jax.jit, static_argnames=["config"]) |
| def unified_attention( |
| q: jax.Array, |
| k: jax.Array, |
| v: jax.Array, |
| temp: jax.Array, |
| segment_ids: jax.Array, |
| config: KernelConfig = KernelConfig(), |
| ) -> jax.Array: |
| """ |
| Unified attention with novel ALSC cap function. |
| |
| Args: |
| q: (B, Q, H, D) query tensor |
| k: (B, KV, H_kv, D) key tensor |
| v: (B, KV, H_kv, D) value tensor |
| temp: (B, Q) or (B, Q, 1) per-query temperature |
| segment_ids: (B, S) in {1, -1, 0} (history, candidate, padding) |
| config: KernelConfig with cap method, block sizes, backend choice |
| """ |
| B, Q, H, D = q.shape |
| _, KV, H_kv, _ = k.shape |
|
|
| assert H % H_kv == 0, f"{H=} must be divisible by {H_kv=}" |
| assert D % 64 == 0, f"head_dim {D} must be multiple of 64" |
| assert q.dtype == k.dtype == v.dtype |
| assert q.dtype in (jnp.float16, jnp.bfloat16) |
|
|
| bounds = SegmentBounds.from_segment_ids(segment_ids) |
| bound_arr = bounds.to_array() |
|
|
| if temp.ndim == 2: |
| temp = temp[..., None] |
|
|
| if config.backend == "triton": |
| return _triton_attention(q, k, v, temp, bound_arr, config) |
| else: |
| raise NotImplementedError("Mosaic backend: use xrex.attention directly") |
|
|
|
|
| def _triton_attention(q, k, v, temp, bound_arr, config): |
| B, Q, H, D = q.shape |
| _, KV, H_kv, _ = k.shape |
|
|
| block_q = min(config.block_q, Q) |
| block_kv = min(config.block_kv, KV) |
| grid = (pl.cdiv(Q, block_q), B, H) |
|
|
| num_warps = config.num_warps |
| if num_warps is None: |
| num_warps = 4 if D <= 64 else 8 |
|
|
| fwd_kernel = make_triton_forward_kernel(config) |
|
|
| out_shape = jax.ShapeDtypeStruct(q.shape, q.dtype) |
| residual_shapes = [ |
| jax.ShapeDtypeStruct((B, H, Q), jnp.float32), |
| jax.ShapeDtypeStruct((B, H, Q), jnp.float32), |
| ] |
|
|
| in_specs = [ |
| pl.BlockSpec(lambda _, b, h: (b, 0, h, 0), (None, Q, None, D)), |
| pl.BlockSpec(lambda _, b, h: (b, 0, h // (H // H_kv), 0), (None, D, None, KV)), |
| pl.BlockSpec(lambda _, b, h: (b, 0, h // (H // H_kv), 0), (None, KV, None, D)), |
| pl.BlockSpec(lambda _, b, h: (b, 0), (None, Q)), |
| pl.BlockSpec(lambda _, b, h: (b, 0), (None, 4)), |
| ] |
| out_specs = pl.BlockSpec(lambda _, b, h: (b, 0, h, 0), (None, Q, None, D)) |
|
|
| out, (l, m) = pl.pallas_call( |
| fwd_kernel, |
| grid=grid, |
| in_specs=in_specs, |
| out_specs=[out_specs] + [pl.BlockSpec(lambda _, b, h: (b, h, 0), (None, None, Q))] * 2, |
| compiler_params=pltriton.CompilerParams( |
| num_warps=num_warps, num_stages=config.num_stages), |
| out_shape=[out_shape] + residual_shapes, |
| name="unified_attention_fwd", |
| )(q, k.swapaxes(1, 3), v, temp, bound_arr) |
|
|
| return out |
|
|
|
|
| def _unified_fwd(q, k, v, temp, segment_ids, config): |
| B, Q, H, D = q.shape |
| bounds = SegmentBounds.from_segment_ids(segment_ids) |
| bound_arr = bounds.to_array() |
|
|
| if temp.ndim == 2: |
| temp = temp[..., None] |
|
|
| out = _triton_attention(q, k, v, temp, bound_arr, config) |
| return out, (q, k, v, temp, bound_arr, out, config) |
|
|
|
|
| def _unified_bwd(config, res, do): |
| q, k, v, temp, bound_arr, out, _ = res |
| B, Q, H, D = q.shape |
| _, KV, H_kv, _ = k.shape |
|
|
| |
| |
| |
| raise NotImplementedError("Backward pass requires Pallas dispatch") |
|
|
|
|
| unified_attention.defvjp(_unified_fwd, _unified_bwd) |
|
|
|
|
| |
| |
| |
|
|
| @functools.partial(jax.jit, static_argnames=["config"]) |
| def unified_attention_reference( |
| q: jax.Array, |
| k: jax.Array, |
| v: jax.Array, |
| temp: jax.Array, |
| segment_ids: jax.Array, |
| config: KernelConfig = KernelConfig(), |
| ) -> jax.Array: |
| """Reference implementation in pure JAX for verification.""" |
| B, Q, H, D = q.shape |
| _, KV, H_kv, _ = k.shape |
| q_h_per_kv = H // H_kv |
|
|
| q_reshaped = q.reshape(B, Q, H_kv, q_h_per_kv, D) |
| logits = jnp.einsum("bqhkc,bkhc->bqhk", q_reshaped, k).astype(jnp.float32) |
| logits *= config.sm_scale |
|
|
| logits = cap_forward(logits, config.cap_method, config.cap_params) |
|
|
| if temp.ndim == 2: |
| temp = temp[..., None] |
| logits *= temp[..., None, :] |
|
|
| bounds = SegmentBounds.from_segment_ids(segment_ids) |
| q_pos = jnp.arange(Q)[None, :, None, None] |
| kv_pos = jnp.arange(KV)[None, None, None, :] |
|
|
| from .segment_bounds import ranker_mask |
| mask = ranker_mask(q_pos, kv_pos, bounds) |
| logits = jnp.where(mask, logits, -jnp.inf) |
|
|
| weights = jax.nn.softmax(logits, axis=-1).astype(q.dtype) |
| out = jnp.einsum("bqhk,bkhc->bqhkc", weights, v) |
| return out.reshape(B, Q, H, D) |
|
|