From a92f665394775b9f291bd50c8d49d80cbf1e1f15 Mon Sep 17 00:00:00 2001 From: Sagar Chapara Date: Wed, 16 Sep 2026 18:39:11 +0000 Subject: [PATCH] Add custom mask-free fused backward and ring backward attention kernels --- .../kernels/custom_splash_attention.py | 989 +++++++++++++++++- .../splash_attention/ring_attention_kernel.py | 300 ++++-- .../tests/custom_splash_backward_test.py | 471 +++++++++ 3 files changed, 1675 insertions(+), 85 deletions(-) create mode 100644 src/maxdiffusion/tests/custom_splash_backward_test.py diff --git a/src/maxdiffusion/kernels/custom_splash_attention.py b/src/maxdiffusion/kernels/custom_splash_attention.py index 6bd3f3493..9cbd678a0 100644 --- a/src/maxdiffusion/kernels/custom_splash_attention.py +++ b/src/maxdiffusion/kernels/custom_splash_attention.py @@ -31,8 +31,26 @@ NT_DIM_NUMBERS = (((1,), (1,)), ((), ())) +LN2 = float(np.log(2.0)) + + class _BlockSizes: - __slots__ = ("block_q", "block_kv", "block_kv_compute", "block_kv_compute_in") + __slots__ = ( + "block_q", + "block_kv", + "block_kv_compute", + "block_kv_compute_in", + "block_q_dkv", + "block_kv_dkv", + "block_kv_dkv_compute", + "block_kv_dkv_compute_in", + "block_q_dq", + "block_kv_dq", + "block_kv_dq_compute", + "block_kv_dq_compute_in", + "use_fused_bwd_kernel", + "dq_reduction_steps", + ) def __init__( self, @@ -40,11 +58,65 @@ def __init__( block_kv: int, block_kv_compute: int | None = None, block_kv_compute_in: int = 256, + block_q_dkv: int | None = None, + block_kv_dkv: int | None = None, + block_kv_dkv_compute: int | None = None, + block_kv_dkv_compute_in: int | None = None, + block_q_dq: int | None = None, + block_kv_dq: int | None = None, + block_kv_dq_compute: int | None = None, + block_kv_dq_compute_in: int | None = None, + use_fused_bwd_kernel: bool = True, + dq_reduction_steps: int | None = 3, ): self.block_q = block_q self.block_kv = block_kv self.block_kv_compute = block_kv_compute if block_kv_compute is not None else block_kv self.block_kv_compute_in = block_kv_compute_in + self.block_q_dkv = block_q_dkv if block_q_dkv is not None else block_q + self.block_kv_dkv = block_kv_dkv if block_kv_dkv is not None else block_kv + self.block_kv_dkv_compute = ( + block_kv_dkv_compute if block_kv_dkv_compute is not None else self.block_kv_dkv + ) + self.block_kv_dkv_compute_in = ( + block_kv_dkv_compute_in if block_kv_dkv_compute_in is not None else block_kv_compute_in + ) + self.block_q_dq = block_q_dq if block_q_dq is not None else block_q + self.block_kv_dq = block_kv_dq if block_kv_dq is not None else block_kv + self.block_kv_dq_compute = ( + block_kv_dq_compute if block_kv_dq_compute is not None else self.block_kv_dq + ) + self.block_kv_dq_compute_in = ( + block_kv_dq_compute_in if block_kv_dq_compute_in is not None else block_kv_compute_in + ) + self.use_fused_bwd_kernel = use_fused_bwd_kernel + self.dq_reduction_steps = dq_reduction_steps + + def _as_tuple(self): + return ( + self.block_q, + self.block_kv, + self.block_kv_compute, + self.block_kv_compute_in, + self.block_q_dkv, + self.block_kv_dkv, + self.block_kv_dkv_compute, + self.block_kv_dkv_compute_in, + self.block_q_dq, + self.block_kv_dq, + self.block_kv_dq_compute, + self.block_kv_dq_compute_in, + self.use_fused_bwd_kernel, + self.dq_reduction_steps, + ) + + def __eq__(self, other): + if not isinstance(other, _BlockSizes): + return False + return self._as_tuple() == other._as_tuple() + + def __hash__(self): + return hash(self._as_tuple()) # Fixed-m softmax-bound constants. Instead of tracking the online-softmax @@ -89,6 +161,7 @@ def _flash_attention_kernel( fuse_reciprocal: bool = True, use_fixed_m: bool = False, uniform_fixed_m: bool = False, + save_lse: bool = False, ): float32 = jnp.float32 head_dim_v_repeats, rem = divmod(head_dim_v, NUM_SUBLANES) @@ -315,7 +388,11 @@ def end(): # can merge shard contributions and normalize only once at the very end. o_ref[...] = o_scratch_ref[...].astype(o_ref.dtype) if l_ring_ref is not None: - l_ring_ref[...] = l.astype(l_ring_ref.dtype) + if save_lse: + log = jnp.log2 if use_base2_exp else jnp.log + l_ring_ref[...] = (m_scratch_ref[...] + log(l)).astype(l_ring_ref.dtype) + else: + l_ring_ref[...] = l.astype(l_ring_ref.dtype) if m_ring_ref is not None: m_ring_ref[...] = m_scratch_ref[...].astype(m_ring_ref.dtype) @@ -483,6 +560,7 @@ def _splash_attention_forward( vmem_limit_bytes: int | None = None, use_fixed_m: bool = False, mk: jax.Array | None = None, + save_residuals: bool = False, ): num_q_heads, padded_q_seq_len, head_dim_qk = q.shape head_dim_v = v.shape[-1] @@ -513,6 +591,15 @@ def k_index_map(h, i, j, *_): def v_index_map(h, i, j, *_): return (h // q_heads_per_kv_head, j, 0) + grid_width = (actual_kv_seq_len + bkv - 1) // bkv + grid_height = (actual_q_seq_len + bq - 1) // bq + active_q_len = grid_height * bq + active_kv_len = grid_width * bkv + + q_in = jnp.pad(q, ((0, 0), (0, active_q_len - q.shape[1]), (0, 0))) if q.shape[1] < active_q_len else q + k_in = jnp.pad(k, ((0, 0), (0, active_kv_len - k.shape[1]), (0, 0))) if k.shape[1] < active_kv_len else k + v_in = jnp.pad(v, ((0, 0), (0, active_kv_len - v.shape[1]), (0, 0))) if v.shape[1] < active_kv_len else v + in_specs = [ pl.BlockSpec((None, bq, head_dim_qk), q_index_map), pl.BlockSpec((None, bkv, head_dim_qk), k_index_map), @@ -522,7 +609,7 @@ def v_index_map(h, i, j, *_): jax.ShapeDtypeStruct((NUM_SUBLANES, bq), jnp.float32), jax.ShapeDtypeStruct((NUM_SUBLANES, bq), jnp.float32), jax.ShapeDtypeStruct((head_dim_v, bq), jnp.float32), - jax.ShapeDtypeStruct((num_q_heads, head_dim_v, actual_q_seq_len), q.dtype), + jax.ShapeDtypeStruct((num_q_heads, head_dim_v, active_q_len), q.dtype), ] out_specs = [ pl.BlockSpec((NUM_SUBLANES, bq), lambda *_: (0, 0)), @@ -530,8 +617,10 @@ def v_index_map(h, i, j, *_): pl.BlockSpec((head_dim_v, bq), lambda *_: (0, 0)), pl.BlockSpec((None, head_dim_v, bq), out_index_map), ] - grid_width = (actual_kv_seq_len + bkv - 1) // bkv - grid_height = (actual_q_seq_len + bq - 1) // bq + if save_residuals: + out_shapes.append(jax.ShapeDtypeStruct((num_q_heads, NUM_SUBLANES, active_q_len), jnp.float32)) + out_specs.append(pl.BlockSpec((None, NUM_SUBLANES, bq), out_index_map)) + grid = (num_q_heads, grid_height, grid_width) all_out = pl.pallas_call( @@ -546,6 +635,7 @@ def v_index_map(h, i, j, *_): kv_seq_len=actual_kv_seq_len, use_base2_exp=use_base2_exp, use_fixed_m=use_fixed_m, + save_lse=save_residuals, ), grid_spec=pltpu.PrefetchScalarGridSpec( num_scalar_prefetch=1, @@ -561,8 +651,11 @@ def v_index_map(h, i, j, *_): vmem_limit_bytes=vmem_limit_bytes, ), out_shape=out_shapes, - )(mk, q, k, v) - return all_out[-1] + )(mk, q_in, k_in, v_in) + out = all_out[3][:, :, :actual_q_seq_len] + if save_residuals: + return out, all_out[4][:, :, :actual_q_seq_len] + return out def _splash_attention_forward_ring( @@ -617,6 +710,15 @@ def k_index_map(h, i, j, *_): def v_index_map(h, i, j, *_): return (h // q_heads_per_kv_head, j, 0) + grid_width = (actual_kv_seq_len + bkv - 1) // bkv + grid_height = (actual_q_seq_len + bq - 1) // bq + active_q_len = grid_height * bq + active_kv_len = grid_width * bkv + + q_in = jnp.pad(q, ((0, 0), (0, active_q_len - q.shape[1]), (0, 0))) if q.shape[1] < active_q_len else q + k_in = jnp.pad(k, ((0, 0), (0, active_kv_len - k.shape[1]), (0, 0))) if k.shape[1] < active_kv_len else k + v_in = jnp.pad(v, ((0, 0), (0, active_kv_len - v.shape[1]), (0, 0))) if v.shape[1] < active_kv_len else v + in_specs = [ pl.BlockSpec((None, bq, head_dim_qk), q_index_map), pl.BlockSpec((None, bkv, head_dim_qk), k_index_map), @@ -626,9 +728,9 @@ def v_index_map(h, i, j, *_): jax.ShapeDtypeStruct((NUM_SUBLANES, bq), jnp.float32), jax.ShapeDtypeStruct((NUM_SUBLANES, bq), jnp.float32), jax.ShapeDtypeStruct((head_dim_v, bq), jnp.float32), - jax.ShapeDtypeStruct((num_q_heads, head_dim_v, actual_q_seq_len), jnp.float32), - jax.ShapeDtypeStruct((num_q_heads, NUM_SUBLANES, actual_q_seq_len), jnp.float32), - jax.ShapeDtypeStruct((num_q_heads, NUM_SUBLANES, actual_q_seq_len), jnp.float32), + jax.ShapeDtypeStruct((num_q_heads, head_dim_v, active_q_len), jnp.float32), + jax.ShapeDtypeStruct((num_q_heads, NUM_SUBLANES, active_q_len), jnp.float32), + jax.ShapeDtypeStruct((num_q_heads, NUM_SUBLANES, active_q_len), jnp.float32), ] out_specs = [ pl.BlockSpec((NUM_SUBLANES, bq), lambda *_: (0, 0)), @@ -638,8 +740,6 @@ def v_index_map(h, i, j, *_): pl.BlockSpec((None, NUM_SUBLANES, bq), out_index_map), pl.BlockSpec((None, NUM_SUBLANES, bq), out_index_map), ] - grid_width = (actual_kv_seq_len + bkv - 1) // bkv - grid_height = (actual_q_seq_len + bq - 1) // bq grid = (num_q_heads, grid_height, grid_width) # Scalar-prefetch operand carrying per-head fixed-m data (same convention as @@ -678,10 +778,10 @@ def v_index_map(h, i, j, *_): vmem_limit_bytes=vmem_limit_bytes, ), out_shape=out_shapes, - )(mk, q, k, v) - out = jnp.swapaxes(all_out[3], 1, 2) # (h, head_dim_v, s) -> (h, s, head_dim_v) - l = all_out[4][:, 0, :] # (h, s) - m = all_out[5][:, 0, :] # (h, s) + )(mk, q_in, k_in, v_in) + out = jnp.swapaxes(all_out[3][:, :, :actual_q_seq_len], 1, 2) # (h, head_dim_v, s) -> (h, s, head_dim_v) + l = all_out[4][:, 0, :actual_q_seq_len] # (h, s) + m = all_out[5][:, 0, :actual_q_seq_len] # (h, s) return out, m, l @@ -774,6 +874,843 @@ def out_index_map(h, i, j, *_): return all_out[-1] +def _flash_attention_dq_kernel( + q_ref, + k_ref, + v_ref, + do_ref, + lse_ref, + di_ref, + dq_scratch_ref, + dq_ref, + *, + grid_width: int, + bkv: int, + bkv_compute: int, + bkv_compute_in: int, + kv_seq_len: int, + use_base2_exp: bool = True, +): + float32 = jnp.float32 + _, _, j = pl.program_id(0), pl.program_id(1), pl.program_id(2) + exp = jnp.exp2 if use_base2_exp else jnp.exp + + @pl.when(j == 0) + def init(): + dq_scratch_ref[...] = jnp.zeros_like(dq_scratch_ref) + + def _dq_inner(qk, k_chunk, v_chunk, do, lse, di, dq_prev): + step = bkv_compute_in + for idx in range(0, qk.shape[0], step): + qk_slice = qk[idx : idx + step] + v_slice = v_chunk[idx : idx + step] + k_slice = k_chunk[idx : idx + step] + + p_curr = exp(qk_slice - lse) + dp_curr = lax.dot_general( + v_slice, + do.astype(v_slice.dtype), + (((1,), (0,)), ((), ())), + preferred_element_type=float32, + ) + ds_curr = p_curr * (dp_curr - di) + dq_curr = lax.dot_general( + ds_curr.astype(k_slice.dtype), + k_slice, + (((0,), (0,)), ((), ())), + preferred_element_type=float32, + ) + dq_prev = dq_prev + dq_curr + return dq_prev + + def compute_body(kv_compute_index, _): + q = q_ref[...] + do = do_ref[...] + lse = lse_ref[0:1, :] + di = di_ref[0:1, :] + base_offset = kv_compute_index * bkv_compute + slice_k = pl.ds(base_offset, bkv_compute) + k_chunk = k_ref[slice_k, :] + v_chunk = v_ref[slice_k, :] + qk = lax.dot_general(k_chunk, q, NT_DIM_NUMBERS, preferred_element_type=float32) + dq_scratch_ref[...] = _dq_inner(qk, k_chunk, v_chunk, do, lse, di, dq_scratch_ref[...]) + + def last_compute_body(kv_compute_index): + q = q_ref[...] + do = do_ref[...] + lse = lse_ref[0:1, :] + di = di_ref[0:1, :] + slice_k_len = kv_seq_len % bkv_compute + slice_k = pl.ds(kv_compute_index * bkv_compute, slice_k_len) + k_chunk = k_ref[slice_k, :] + v_chunk = v_ref[slice_k, :] + qk = lax.dot_general(k_chunk, q, NT_DIM_NUMBERS, preferred_element_type=float32) + dq_scratch_ref[...] = _dq_inner(qk, k_chunk, v_chunk, do, lse, di, dq_scratch_ref[...]) + + assert bkv % bkv_compute == 0 + + @pl.when(j != grid_width - 1) + def body(): + lax.fori_loop(0, (bkv // bkv_compute), compute_body, None, unroll=True) + + @pl.when(j == grid_width - 1) + def last_body(): + if kv_seq_len % bkv == 0: + iter_num = bkv // bkv_compute + lax.fori_loop(0, iter_num, compute_body, None, unroll=True) + else: + remain_kv_seq_len = kv_seq_len % bkv + iter_num = (remain_kv_seq_len + bkv_compute - 1) // bkv_compute + if remain_kv_seq_len % bkv_compute == 0: + lax.fori_loop(0, iter_num, compute_body, None, unroll=True) + else: + lax.fori_loop(0, iter_num - 1, compute_body, None, unroll=True) + last_compute_body(iter_num - 1) + + @pl.when(j == grid_width - 1) + def end(): + if use_base2_exp: + dq_ref[...] = (dq_scratch_ref[...] * LN2).astype(dq_ref.dtype) + else: + dq_ref[...] = dq_scratch_ref[...].astype(dq_ref.dtype) + + +def _flash_attention_dkv_kernel( + q_ref, + k_ref, + v_ref, + do_ref, + lse_ref, + di_ref, + dk_scratch_ref, + dv_scratch_ref, + dk_ref, + dv_ref, + *, + grid_width: int, + total_q_steps: int, + bkv: int, + bkv_compute: int, + bkv_compute_in: int, + kv_seq_len: int, + use_base2_exp: bool = True, +): + float32 = jnp.float32 + _, j, step_q = pl.program_id(0), pl.program_id(1), pl.program_id(2) + exp = jnp.exp2 if use_base2_exp else jnp.exp + + @pl.when(step_q == 0) + def init(): + dk_scratch_ref[...] = jnp.zeros_like(dk_scratch_ref) + dv_scratch_ref[...] = jnp.zeros_like(dv_scratch_ref) + + def _dkv_inner(base_offset, qk, q, v_chunk, do, lse, di): + step = bkv_compute_in + for idx in range(0, qk.shape[0], step): + sub_len = min(step, qk.shape[0] - idx) + sub_slice = pl.ds(base_offset + idx, sub_len) + qk_slice = qk[idx : idx + sub_len] + v_slice = v_chunk[idx : idx + sub_len] + + p_curr = exp(qk_slice - lse) + dv_curr = lax.dot_general( + p_curr.astype(do.dtype), + do, + NT_DIM_NUMBERS, + preferred_element_type=float32, + ) + dv_scratch_ref[sub_slice, :] = dv_scratch_ref[sub_slice, :] + dv_curr + + dp_curr = lax.dot_general( + v_slice, + do.astype(v_slice.dtype), + (((1,), (0,)), ((), ())), + preferred_element_type=float32, + ) + ds_curr = p_curr * (dp_curr - di) + dk_curr = lax.dot_general( + ds_curr.astype(q.dtype), + q, + (((1,), (0,)), ((), ())), + preferred_element_type=float32, + ) + dk_scratch_ref[sub_slice, :] = dk_scratch_ref[sub_slice, :] + dk_curr + + def compute_body(kv_compute_index, _): + q = q_ref[...] + do = do_ref[...] + lse = lse_ref[0:1, :] + di = di_ref[0:1, :] + base_offset = kv_compute_index * bkv_compute + slice_k = pl.ds(base_offset, bkv_compute) + k_chunk = k_ref[slice_k, :] + v_chunk = v_ref[slice_k, :] + qk = lax.dot_general(k_chunk, q, NT_DIM_NUMBERS, preferred_element_type=float32) + _dkv_inner(base_offset, qk, q, v_chunk, do, lse, di) + + def last_compute_body(kv_compute_index): + q = q_ref[...] + do = do_ref[...] + lse = lse_ref[0:1, :] + di = di_ref[0:1, :] + base_offset = kv_compute_index * bkv_compute + slice_k_len = kv_seq_len % bkv_compute + slice_k = pl.ds(base_offset, slice_k_len) + k_chunk = k_ref[slice_k, :] + v_chunk = v_ref[slice_k, :] + qk = lax.dot_general(k_chunk, q, NT_DIM_NUMBERS, preferred_element_type=float32) + _dkv_inner(base_offset, qk, q, v_chunk, do, lse, di) + + assert bkv % bkv_compute == 0 + + @pl.when(j != grid_width - 1) + def body(): + lax.fori_loop(0, (bkv // bkv_compute), compute_body, None, unroll=True) + + @pl.when(j == grid_width - 1) + def last_body(): + if kv_seq_len % bkv == 0: + iter_num = bkv // bkv_compute + lax.fori_loop(0, iter_num, compute_body, None, unroll=True) + else: + remain_kv_seq_len = kv_seq_len % bkv + iter_num = (remain_kv_seq_len + bkv_compute - 1) // bkv_compute + if remain_kv_seq_len % bkv_compute == 0: + lax.fori_loop(0, iter_num, compute_body, None, unroll=True) + else: + lax.fori_loop(0, iter_num - 1, compute_body, None, unroll=True) + last_compute_body(iter_num - 1) + + @pl.when(step_q == total_q_steps - 1) + def end(): + if use_base2_exp: + dk_ref[...] = (dk_scratch_ref[...] * LN2).astype(dk_ref.dtype) + else: + dk_ref[...] = dk_scratch_ref[...].astype(dk_ref.dtype) + dv_ref[...] = dv_scratch_ref[...].astype(dv_ref.dtype) + + +def _flash_attention_bwd_fused_kernel( + q_ref, + k_ref, + v_ref, + do_ref, + lse_ref, + di_ref, + dq_alias_ref, + dq_ref, + dk_ref, + dv_ref, + dq_scratch_ref, + dk_scratch_ref, + dv_scratch_ref, + *, + grid_width: int, + grid_height: int, + q_heads_per_kv_head: int, + bkv: int, + bkv_compute: int, + bkv_compute_in: int, + kv_seq_len: int, + use_base2_exp: bool = True, + use_dq_aliasing: bool = False, +): + """Fused backward kernel iterating over KV outer (grid_width) and Q inner (num_q_heads, grid_height). + + Computes dQ, dK, and dV in a single pass without masks, ignoring KV tail padding via + slice bounds and accumulating dK/dV in VMEM across Q steps and KV head groups. + """ + float32 = jnp.float32 + j, h_q, i = pl.program_id(0), pl.program_id(1), pl.program_id(2) + exp = jnp.exp2 if use_base2_exp else jnp.exp + + q_head_in_group = lax.rem(h_q, q_heads_per_kv_head) + should_init_dkv = jnp.logical_and(i == 0, q_head_in_group == 0) + should_write_dkv = jnp.logical_and( + i == grid_height - 1, q_head_in_group == q_heads_per_kv_head - 1 + ) + + @pl.when(should_init_dkv) + def init_dkv(): + dk_scratch_ref[...] = jnp.zeros_like(dk_scratch_ref) + dv_scratch_ref[...] = jnp.zeros_like(dv_scratch_ref) + + dq_scratch_ref[...] = jnp.zeros_like(dq_scratch_ref) + + def _bwd_fused_inner(base_offset, qk, k_chunk, v_chunk, q, do, lse, di): + step = bkv_compute_in + dq_acc = dq_scratch_ref[...] + for idx in range(0, qk.shape[0], step): + sub_len = min(step, qk.shape[0] - idx) + sub_slice = pl.ds(base_offset + idx, sub_len) + qk_slice = qk[idx : idx + sub_len] + k_slice = k_chunk[idx : idx + sub_len] + v_slice = v_chunk[idx : idx + sub_len] + + p_curr = exp(qk_slice - lse) + dv_curr = lax.dot_general( + p_curr.astype(do.dtype), + do, + NT_DIM_NUMBERS, + preferred_element_type=float32, + ) + dv_scratch_ref[sub_slice, :] = dv_scratch_ref[sub_slice, :] + dv_curr + + dp_curr = lax.dot_general( + v_slice, + do.astype(v_slice.dtype), + (((1,), (0,)), ((), ())), + preferred_element_type=float32, + ) + ds_curr = p_curr * (dp_curr - di) + dk_curr = lax.dot_general( + ds_curr.astype(q.dtype), + q, + (((1,), (0,)), ((), ())), + preferred_element_type=float32, + ) + dk_scratch_ref[sub_slice, :] = dk_scratch_ref[sub_slice, :] + dk_curr + + dq_curr = lax.dot_general( + ds_curr.astype(k_slice.dtype), + k_slice, + (((0,), (0,)), ((), ())), + preferred_element_type=float32, + ) + dq_acc = dq_acc + dq_curr + dq_scratch_ref[...] = dq_acc + + def compute_body(kv_compute_index, _): + q = q_ref[...] + do = do_ref[...] + lse = lse_ref[0:1, :] + di = di_ref[0:1, :] + base_offset = kv_compute_index * bkv_compute + slice_k = pl.ds(base_offset, bkv_compute) + k_chunk = k_ref[slice_k, :] + v_chunk = v_ref[slice_k, :] + qk = lax.dot_general(k_chunk, q, NT_DIM_NUMBERS, preferred_element_type=float32) + _bwd_fused_inner(base_offset, qk, k_chunk, v_chunk, q, do, lse, di) + + def last_compute_body(kv_compute_index): + q = q_ref[...] + do = do_ref[...] + lse = lse_ref[0:1, :] + di = di_ref[0:1, :] + base_offset = kv_compute_index * bkv_compute + slice_k_len = kv_seq_len % bkv_compute + slice_k = pl.ds(base_offset, slice_k_len) + k_chunk = k_ref[slice_k, :] + v_chunk = v_ref[slice_k, :] + qk = lax.dot_general(k_chunk, q, NT_DIM_NUMBERS, preferred_element_type=float32) + _bwd_fused_inner(base_offset, qk, k_chunk, v_chunk, q, do, lse, di) + + assert bkv % bkv_compute == 0 + + @pl.when(j != grid_width - 1) + def body(): + lax.fori_loop(0, (bkv // bkv_compute), compute_body, None, unroll=True) + + @pl.when(j == grid_width - 1) + def last_body(): + if kv_seq_len % bkv == 0: + iter_num = bkv // bkv_compute + lax.fori_loop(0, iter_num, compute_body, None, unroll=True) + else: + remain_kv_seq_len = kv_seq_len % bkv + iter_num = (remain_kv_seq_len + bkv_compute - 1) // bkv_compute + if remain_kv_seq_len % bkv_compute == 0: + lax.fori_loop(0, iter_num, compute_body, None, unroll=True) + else: + lax.fori_loop(0, iter_num - 1, compute_body, None, unroll=True) + last_compute_body(iter_num - 1) + + if use_base2_exp: + dq_val = (dq_scratch_ref[...] * LN2).astype(dq_ref.dtype) + else: + dq_val = dq_scratch_ref[...].astype(dq_ref.dtype) + + if use_dq_aliasing: + + @pl.when(j < 3) + def write_dq_first(): + dq_ref[...] = dq_val + + @pl.when(j >= 3) + def write_dq_acc(): + dq_ref[...] = dq_alias_ref[...] + dq_val + + else: + dq_ref[...] = dq_val + + @pl.when(should_write_dkv) + def end_dkv(): + if use_base2_exp: + dk_ref[...] = (dk_scratch_ref[...] * LN2).astype(dk_ref.dtype) + else: + dk_ref[...] = dk_scratch_ref[...].astype(dk_ref.dtype) + dv_ref[...] = dv_scratch_ref[...].astype(dv_ref.dtype) + + +def _splash_attention_bwd_fused( + q: jax.Array, + k: jax.Array, + v: jax.Array, + do: jax.Array, + lse: jax.Array, + di: jax.Array, + block_sizes: _BlockSizes, + actual_q_seq_len: int, + actual_kv_seq_len: int, + padded_q_seq_len: int, + padded_kv_seq_len: int, + use_base2_exp: bool = True, + use_experimental_scheduler: bool = False, + vmem_limit_bytes: int | None = None, +): + num_q_heads, _, head_dim_qk = q.shape + head_dim_v = v.shape[-1] + num_kv_heads = k.shape[0] + q_heads_per_kv_head = num_q_heads // num_kv_heads + + bq = block_sizes.block_q_dkv + bkv = block_sizes.block_kv_dkv + bkv_compute = block_sizes.block_kv_dkv_compute + bkv_compute_in = block_sizes.block_kv_dkv_compute_in + + grid_width = (actual_kv_seq_len + bkv - 1) // bkv + grid_height = (actual_q_seq_len + bq - 1) // bq + active_q_len = grid_height * bq + active_kv_len = grid_width * bkv + + if actual_q_seq_len < active_q_len: + pad_q = active_q_len - actual_q_seq_len + q_bwd = jnp.pad(q[:, :actual_q_seq_len, :], ((0, 0), (0, pad_q), (0, 0))) + do_bwd = jnp.pad(do[:, :, :actual_q_seq_len], ((0, 0), (0, 0), (0, pad_q))) + lse_bwd = jnp.pad(lse[:, :, :actual_q_seq_len], ((0, 0), (0, 0), (0, pad_q))) + di_bwd = jnp.pad(di[:, :, :actual_q_seq_len], ((0, 0), (0, 0), (0, pad_q))) + elif q.shape[1] > active_q_len: + q_bwd = q[:, :active_q_len, :] + do_bwd = do[:, :, :active_q_len] + lse_bwd = lse[:, :, :active_q_len] + di_bwd = di[:, :, :active_q_len] + else: + q_bwd, do_bwd, lse_bwd, di_bwd = q, do, lse, di + + k_bwd = jnp.pad(k, ((0, 0), (0, active_kv_len - k.shape[1]), (0, 0))) if k.shape[1] < active_kv_len else k + v_bwd = jnp.pad(v, ((0, 0), (0, active_kv_len - v.shape[1]), (0, 0))) if v.shape[1] < active_kv_len else v + + use_dq_aliasing = ( + block_sizes.dq_reduction_steps == 3 and grid_width > 3 + ) + + if use_dq_aliasing: + dq_index_map = lambda j, h_q, i, *_: (j % 3, h_q, i, 0) + dq_spec = pl.BlockSpec((None, None, bq, head_dim_qk), dq_index_map) + dq_alias_spec = dq_spec + dq_dtype = jnp.float32 + dq_shape = jax.ShapeDtypeStruct((3, num_q_heads, active_q_len, head_dim_qk), dq_dtype) + dq_init = lax.empty((3, num_q_heads, active_q_len, head_dim_qk), dtype=dq_dtype) + else: + dq_index_map = lambda j, h_q, i, *_: (j, h_q, i, 0) + dq_spec = pl.BlockSpec((None, None, bq, head_dim_qk), dq_index_map) + dq_alias_spec = None + dq_dtype = q.dtype if grid_width == 1 else jnp.float32 + dq_shape = jax.ShapeDtypeStruct((grid_width, num_q_heads, active_q_len, head_dim_qk), dq_dtype) + dq_init = None + + in_specs = [ + pl.BlockSpec((None, bq, head_dim_qk), lambda j, h_q, i, *_: (h_q, i, 0)), + pl.BlockSpec((None, bkv, head_dim_qk), lambda j, h_q, i, *_: (h_q // q_heads_per_kv_head, j, 0)), + pl.BlockSpec((None, bkv, head_dim_v), lambda j, h_q, i, *_: (h_q // q_heads_per_kv_head, j, 0)), + pl.BlockSpec((None, head_dim_v, bq), lambda j, h_q, i, *_: (h_q, 0, i)), + pl.BlockSpec((None, NUM_SUBLANES, bq), lambda j, h_q, i, *_: (h_q, 0, i)), + pl.BlockSpec((None, NUM_SUBLANES, bq), lambda j, h_q, i, *_: (h_q, 0, i)), + dq_alias_spec, + ] + + out_shapes = [ + dq_shape, + jax.ShapeDtypeStruct((num_kv_heads, active_kv_len, head_dim_qk), k.dtype), + jax.ShapeDtypeStruct((num_kv_heads, active_kv_len, head_dim_v), v.dtype), + jax.ShapeDtypeStruct((bq, head_dim_qk), jnp.float32), + jax.ShapeDtypeStruct((bkv, head_dim_qk), jnp.float32), + jax.ShapeDtypeStruct((bkv, head_dim_v), jnp.float32), + ] + out_specs = [ + dq_spec, + pl.BlockSpec((None, bkv, head_dim_qk), lambda j, h_q, i, *_: (h_q // q_heads_per_kv_head, j, 0)), + pl.BlockSpec((None, bkv, head_dim_v), lambda j, h_q, i, *_: (h_q // q_heads_per_kv_head, j, 0)), + pl.BlockSpec((bq, head_dim_qk), lambda *_: (0, 0)), + pl.BlockSpec((bkv, head_dim_qk), lambda *_: (0, 0)), + pl.BlockSpec((bkv, head_dim_v), lambda *_: (0, 0)), + ] + grid = (grid_width, num_q_heads, grid_height) + input_output_aliases = {6: 0} if use_dq_aliasing else {} + + all_out = pl.pallas_call( + functools.partial( + _flash_attention_bwd_fused_kernel, + grid_width=grid_width, + grid_height=grid_height, + q_heads_per_kv_head=q_heads_per_kv_head, + bkv=bkv, + bkv_compute=bkv_compute, + bkv_compute_in=bkv_compute_in, + kv_seq_len=actual_kv_seq_len, + use_base2_exp=use_base2_exp, + use_dq_aliasing=use_dq_aliasing, + ), + grid_spec=pltpu.PrefetchScalarGridSpec( + num_scalar_prefetch=0, + in_specs=in_specs, + out_specs=out_specs, + grid=grid, + ), + compiler_params=pltpu.CompilerParams( + dimension_semantics=("arbitrary", "arbitrary", "arbitrary"), + flags={"XLA_TPU_FORCE_LP_LLO_SCHEDULER": use_experimental_scheduler}, + disable_bounds_checks=True, + skip_device_barrier=True, + vmem_limit_bytes=vmem_limit_bytes, + ), + out_shape=out_shapes, + input_output_aliases=input_output_aliases, + )(q_bwd, k_bwd, v_bwd, do_bwd, lse_bwd, di_bwd, dq_init) + + dq_unreduced, dk, dv = all_out[0], all_out[1], all_out[2] + if grid_width == 1: + dq = dq_unreduced[0].astype(q.dtype) + else: + dq = dq_unreduced.sum(axis=0).astype(q.dtype) + + if active_q_len > padded_q_seq_len: + dq = dq[:, :padded_q_seq_len, :] + elif active_q_len < padded_q_seq_len: + dq = jnp.pad(dq, ((0, 0), (0, padded_q_seq_len - active_q_len), (0, 0))) + + if active_kv_len > padded_kv_seq_len: + dk = dk[:, :padded_kv_seq_len, :] + dv = dv[:, :padded_kv_seq_len, :] + elif active_kv_len < padded_kv_seq_len: + pad_kv = padded_kv_seq_len - active_kv_len + dk = jnp.pad(dk, ((0, 0), (0, pad_kv), (0, 0))) + dv = jnp.pad(dv, ((0, 0), (0, pad_kv), (0, 0))) + + return dq, dk, dv + + +def _splash_attention_backward( + q: jax.Array, + k: jax.Array, + v: jax.Array, + o: jax.Array, + lse: jax.Array, + do: jax.Array, + block_sizes: _BlockSizes, + q_seq_len: int | None = None, + kv_seq_len: int | None = None, + use_base2_exp: bool = True, + use_experimental_scheduler: bool = False, + vmem_limit_bytes: int | None = None, + di: jax.Array | None = None, +): + num_q_heads, padded_q_seq_len, head_dim_qk = q.shape + head_dim_v = v.shape[-1] + num_kv_heads = k.shape[0] + padded_kv_seq_len = k.shape[1] + + actual_q_seq_len = q_seq_len if q_seq_len is not None else padded_q_seq_len + actual_kv_seq_len = kv_seq_len if kv_seq_len is not None else padded_kv_seq_len + q_heads_per_kv_head = num_q_heads // num_kv_heads + + if di is None: + di_vec = jnp.sum(o.astype(jnp.float32) * do.astype(jnp.float32), axis=1) + di = jnp.broadcast_to(di_vec[:, None, :], (num_q_heads, NUM_SUBLANES, o.shape[2])) + + if block_sizes.use_fused_bwd_kernel: + return _splash_attention_bwd_fused( + q=q, + k=k, + v=v, + do=do, + lse=lse, + di=di, + block_sizes=block_sizes, + actual_q_seq_len=actual_q_seq_len, + actual_kv_seq_len=actual_kv_seq_len, + padded_q_seq_len=padded_q_seq_len, + padded_kv_seq_len=padded_kv_seq_len, + use_base2_exp=use_base2_exp, + use_experimental_scheduler=use_experimental_scheduler, + vmem_limit_bytes=vmem_limit_bytes, + ) + + # --- 1. Compute dQ --- + bq_dq, bkv_dq = block_sizes.block_q_dq, block_sizes.block_kv_dq + bkv_dq_compute = block_sizes.block_kv_dq_compute + bkv_dq_compute_in = block_sizes.block_kv_dq_compute_in + grid_width_dq = (actual_kv_seq_len + bkv_dq - 1) // bkv_dq + grid_height_dq = (actual_q_seq_len + bq_dq - 1) // bq_dq + active_q_len_dq = grid_height_dq * bq_dq + active_kv_len_dq = grid_width_dq * bkv_dq + + if actual_q_seq_len < active_q_len_dq: + pad_q = active_q_len_dq - actual_q_seq_len + q_dq = jnp.pad(q[:, :actual_q_seq_len, :], ((0, 0), (0, pad_q), (0, 0))) + do_dq = jnp.pad(do[:, :, :actual_q_seq_len], ((0, 0), (0, 0), (0, pad_q))) + lse_dq = jnp.pad(lse[:, :, :actual_q_seq_len], ((0, 0), (0, 0), (0, pad_q))) + di_dq = jnp.pad(di[:, :, :actual_q_seq_len], ((0, 0), (0, 0), (0, pad_q))) + elif q.shape[1] > active_q_len_dq: + q_dq = q[:, :active_q_len_dq, :] + do_dq = do[:, :, :active_q_len_dq] + lse_dq = lse[:, :, :active_q_len_dq] + di_dq = di[:, :, :active_q_len_dq] + else: + q_dq, do_dq, lse_dq, di_dq = q, do, lse, di + + k_dq = jnp.pad(k, ((0, 0), (0, active_kv_len_dq - k.shape[1]), (0, 0))) if k.shape[1] < active_kv_len_dq else k + v_dq = jnp.pad(v, ((0, 0), (0, active_kv_len_dq - v.shape[1]), (0, 0))) if v.shape[1] < active_kv_len_dq else v + + dq_in_specs = [ + pl.BlockSpec((None, bq_dq, head_dim_qk), lambda h, i, j, *_: (h, i, 0)), + pl.BlockSpec((None, bkv_dq, head_dim_qk), lambda h, i, j, *_: (h // q_heads_per_kv_head, j, 0)), + pl.BlockSpec((None, bkv_dq, head_dim_v), lambda h, i, j, *_: (h // q_heads_per_kv_head, j, 0)), + pl.BlockSpec((None, head_dim_v, bq_dq), lambda h, i, j, *_: (h, 0, i)), + pl.BlockSpec((None, NUM_SUBLANES, bq_dq), lambda h, i, j, *_: (h, 0, i)), + pl.BlockSpec((None, NUM_SUBLANES, bq_dq), lambda h, i, j, *_: (h, 0, i)), + ] + dq_out_shapes = [ + jax.ShapeDtypeStruct((bq_dq, head_dim_qk), jnp.float32), + jax.ShapeDtypeStruct((num_q_heads, active_q_len_dq, head_dim_qk), q.dtype), + ] + dq_out_specs = [ + pl.BlockSpec((bq_dq, head_dim_qk), lambda *_: (0, 0)), + pl.BlockSpec((None, bq_dq, head_dim_qk), lambda h, i, j, *_: (h, i, 0)), + ] + dq_grid = (num_q_heads, grid_height_dq, grid_width_dq) + + _, dq = pl.pallas_call( + functools.partial( + _flash_attention_dq_kernel, + grid_width=grid_width_dq, + bkv=bkv_dq, + bkv_compute=bkv_dq_compute, + bkv_compute_in=bkv_dq_compute_in, + kv_seq_len=actual_kv_seq_len, + use_base2_exp=use_base2_exp, + ), + grid_spec=pltpu.PrefetchScalarGridSpec( + num_scalar_prefetch=0, + in_specs=dq_in_specs, + out_specs=dq_out_specs, + grid=dq_grid, + ), + compiler_params=pltpu.CompilerParams( + dimension_semantics=("parallel", "arbitrary", "arbitrary"), + flags={"XLA_TPU_FORCE_LP_LLO_SCHEDULER": use_experimental_scheduler}, + disable_bounds_checks=True, + skip_device_barrier=True, + vmem_limit_bytes=vmem_limit_bytes, + ), + out_shape=dq_out_shapes, + )(q_dq, k_dq, v_dq, do_dq, lse_dq, di_dq) + + if active_q_len_dq > padded_q_seq_len: + dq = dq[:, :padded_q_seq_len, :] + elif active_q_len_dq < padded_q_seq_len: + dq = jnp.pad(dq, ((0, 0), (0, padded_q_seq_len - active_q_len_dq), (0, 0))) + + # --- 2. Compute dK, dV --- + bq_dkv, bkv_dkv = block_sizes.block_q_dkv, block_sizes.block_kv_dkv + bkv_dkv_compute = block_sizes.block_kv_dkv_compute + bkv_dkv_compute_in = block_sizes.block_kv_dkv_compute_in + grid_width_dkv = (actual_kv_seq_len + bkv_dkv - 1) // bkv_dkv + grid_height_dkv = (actual_q_seq_len + bq_dkv - 1) // bq_dkv + active_q_len_dkv = grid_height_dkv * bq_dkv + active_kv_len_dkv = grid_width_dkv * bkv_dkv + total_q_steps = q_heads_per_kv_head * grid_height_dkv + + if actual_q_seq_len < active_q_len_dkv: + pad_q = active_q_len_dkv - actual_q_seq_len + q_dkv = jnp.pad(q[:, :actual_q_seq_len, :], ((0, 0), (0, pad_q), (0, 0))) + do_dkv = jnp.pad(do[:, :, :actual_q_seq_len], ((0, 0), (0, 0), (0, pad_q))) + lse_dkv = jnp.pad(lse[:, :, :actual_q_seq_len], ((0, 0), (0, 0), (0, pad_q))) + di_dkv = jnp.pad(di[:, :, :actual_q_seq_len], ((0, 0), (0, 0), (0, pad_q))) + elif q.shape[1] > active_q_len_dkv: + q_dkv = q[:, :active_q_len_dkv, :] + do_dkv = do[:, :, :active_q_len_dkv] + lse_dkv = lse[:, :, :active_q_len_dkv] + di_dkv = di[:, :, :active_q_len_dkv] + else: + q_dkv, do_dkv, lse_dkv, di_dkv = q, do, lse, di + + k_dkv = jnp.pad(k, ((0, 0), (0, active_kv_len_dkv - k.shape[1]), (0, 0))) if k.shape[1] < active_kv_len_dkv else k + v_dkv = jnp.pad(v, ((0, 0), (0, active_kv_len_dkv - v.shape[1]), (0, 0))) if v.shape[1] < active_kv_len_dkv else v + + def q_step_map(h_kv, j, step_q, *_): + h_q = h_kv * q_heads_per_kv_head + (step_q // grid_height_dkv) + i = step_q % grid_height_dkv + return (h_q, i, 0) + + def do_step_map(h_kv, j, step_q, *_): + h_q = h_kv * q_heads_per_kv_head + (step_q // grid_height_dkv) + i = step_q % grid_height_dkv + return (h_q, 0, i) + + dkv_in_specs = [ + pl.BlockSpec((None, bq_dkv, head_dim_qk), q_step_map), + pl.BlockSpec((None, bkv_dkv, head_dim_qk), lambda h_kv, j, step_q, *_: (h_kv, j, 0)), + pl.BlockSpec((None, bkv_dkv, head_dim_v), lambda h_kv, j, step_q, *_: (h_kv, j, 0)), + pl.BlockSpec((None, head_dim_v, bq_dkv), do_step_map), + pl.BlockSpec((None, NUM_SUBLANES, bq_dkv), do_step_map), + pl.BlockSpec((None, NUM_SUBLANES, bq_dkv), do_step_map), + ] + dkv_out_shapes = [ + jax.ShapeDtypeStruct((bkv_dkv, head_dim_qk), jnp.float32), + jax.ShapeDtypeStruct((bkv_dkv, head_dim_v), jnp.float32), + jax.ShapeDtypeStruct((num_kv_heads, active_kv_len_dkv, head_dim_qk), k.dtype), + jax.ShapeDtypeStruct((num_kv_heads, active_kv_len_dkv, head_dim_v), v.dtype), + ] + dkv_out_specs = [ + pl.BlockSpec((bkv_dkv, head_dim_qk), lambda *_: (0, 0)), + pl.BlockSpec((bkv_dkv, head_dim_v), lambda *_: (0, 0)), + pl.BlockSpec((None, bkv_dkv, head_dim_qk), lambda h_kv, j, step_q, *_: (h_kv, j, 0)), + pl.BlockSpec((None, bkv_dkv, head_dim_v), lambda h_kv, j, step_q, *_: (h_kv, j, 0)), + ] + dkv_grid = (num_kv_heads, grid_width_dkv, total_q_steps) + + _, _, dk, dv = pl.pallas_call( + functools.partial( + _flash_attention_dkv_kernel, + grid_width=grid_width_dkv, + total_q_steps=total_q_steps, + bkv=bkv_dkv, + bkv_compute=bkv_dkv_compute, + bkv_compute_in=bkv_dkv_compute_in, + kv_seq_len=actual_kv_seq_len, + use_base2_exp=use_base2_exp, + ), + grid_spec=pltpu.PrefetchScalarGridSpec( + num_scalar_prefetch=0, + in_specs=dkv_in_specs, + out_specs=dkv_out_specs, + grid=dkv_grid, + ), + compiler_params=pltpu.CompilerParams( + dimension_semantics=("parallel", "arbitrary", "arbitrary"), + flags={"XLA_TPU_FORCE_LP_LLO_SCHEDULER": use_experimental_scheduler}, + disable_bounds_checks=True, + skip_device_barrier=True, + vmem_limit_bytes=vmem_limit_bytes, + ), + out_shape=dkv_out_shapes, + )(q_dkv, k_dkv, v_dkv, do_dkv, lse_dkv, di_dkv) + + if active_kv_len_dkv > padded_kv_seq_len: + dk = dk[:, :padded_kv_seq_len, :] + dv = dv[:, :padded_kv_seq_len, :] + elif active_kv_len_dkv < padded_kv_seq_len: + pad_kv = padded_kv_seq_len - active_kv_len_dkv + dk = jnp.pad(dk, ((0, 0), (0, pad_kv), (0, 0))) + dv = jnp.pad(dv, ((0, 0), (0, pad_kv), (0, 0))) + + return dq, dk, dv + + +@functools.partial( + jax.custom_vjp, + nondiff_argnames=( + "block_sizes", + "q_seq_len", + "kv_seq_len", + "use_base2_exp", + "use_experimental_scheduler", + "vmem_limit_bytes", + ), +) +def _splash_attention_custom( + q: jax.Array, + k: jax.Array, + v: jax.Array, + block_sizes: _BlockSizes, + q_seq_len: int | None = None, + kv_seq_len: int | None = None, + use_base2_exp: bool = True, + use_experimental_scheduler: bool = False, + vmem_limit_bytes: int | None = None, +): + return _splash_attention_forward( + q, + k, + v, + block_sizes, + q_seq_len=q_seq_len, + kv_seq_len=kv_seq_len, + use_base2_exp=use_base2_exp, + use_experimental_scheduler=use_experimental_scheduler, + vmem_limit_bytes=vmem_limit_bytes, + save_residuals=False, + ) + + +def _splash_attention_fwd( + q: jax.Array, + k: jax.Array, + v: jax.Array, + block_sizes: _BlockSizes, + q_seq_len: int | None = None, + kv_seq_len: int | None = None, + use_base2_exp: bool = True, + use_experimental_scheduler: bool = False, + vmem_limit_bytes: int | None = None, +): + out, lse = _splash_attention_forward( + q, + k, + v, + block_sizes, + q_seq_len=q_seq_len, + kv_seq_len=kv_seq_len, + use_base2_exp=use_base2_exp, + use_experimental_scheduler=use_experimental_scheduler, + vmem_limit_bytes=vmem_limit_bytes, + save_residuals=True, + ) + return out, (q, k, v, out, lse) + + +def _splash_attention_bwd( + block_sizes: _BlockSizes, + q_seq_len: int | None, + kv_seq_len: int | None, + use_base2_exp: bool, + use_experimental_scheduler: bool, + vmem_limit_bytes: int | None, + residuals, + do: jax.Array, +): + q, k, v, out, lse = residuals + dq, dk, dv = _splash_attention_backward( + q, + k, + v, + out, + lse, + do, + block_sizes, + q_seq_len=q_seq_len, + kv_seq_len=kv_seq_len, + use_base2_exp=use_base2_exp, + use_experimental_scheduler=use_experimental_scheduler, + vmem_limit_bytes=vmem_limit_bytes, + ) + return dq, dk, dv + + +_splash_attention_custom.defvjp(_splash_attention_fwd, _splash_attention_bwd) + + def make_splash_mha( block_sizes: _BlockSizes, orig_q_seq_len: int | None = None, @@ -800,18 +1737,30 @@ def _splash_attention(q, k, v, mk=None): use_experimental_scheduler=use_experimental_scheduler, vmem_limit_bytes=vmem_limit_bytes, ) - return _splash_attention_forward( + if use_fixed_m or mk is not None: + return _splash_attention_forward( + q, + k, + v, + block_sizes, + q_seq_len=orig_q_seq_len, + kv_seq_len=orig_kv_seq_len, + use_base2_exp=use_base2_exp, + use_experimental_scheduler=use_experimental_scheduler, + vmem_limit_bytes=vmem_limit_bytes, + use_fixed_m=use_fixed_m, + mk=mk, + ) + return _splash_attention_custom( q, k, v, - block_sizes, + block_sizes=block_sizes, q_seq_len=orig_q_seq_len, kv_seq_len=orig_kv_seq_len, use_base2_exp=use_base2_exp, use_experimental_scheduler=use_experimental_scheduler, vmem_limit_bytes=vmem_limit_bytes, - use_fixed_m=use_fixed_m, - mk=mk, ) return _splash_attention diff --git a/src/maxdiffusion/kernels/splash_attention/ring_attention_kernel.py b/src/maxdiffusion/kernels/splash_attention/ring_attention_kernel.py index bc49c5af7..afd6ee655 100644 --- a/src/maxdiffusion/kernels/splash_attention/ring_attention_kernel.py +++ b/src/maxdiffusion/kernels/splash_attention/ring_attention_kernel.py @@ -859,8 +859,9 @@ def _custom_ring_attention_forward( bidirectional: bool = False, use_fixed_m: bool = False, fixed_m_norms: tuple[jax.Array, jax.Array] | None = None, -) -> jax.Array: - """Forward-only ring attention using the custom dense splash kernel. + save_residuals: bool = False, +) -> jax.Array | tuple[jax.Array, jax.Array]: + """Ring attention using the custom dense splash kernel. Args: q: Query shard, shape `(num_q_heads, q_seq_len, head_dim_qk)`. Stationary @@ -870,7 +871,6 @@ def _custom_ring_attention_forward( the ring axis. v: Value shard, shape `(num_kv_heads, kv_seq_len, head_dim_v)`. Rotated. block_sizes: Custom-kernel block sizes (block_q / block_kv / block_kv_compute). - bkv_compute_in: Inner VPU register-tiling step for the custom kernel. orig_q_seq_len: Un-padded local query length (grid bound). orig_kv_seq_len: Un-padded local key/value length (grid bound). Assumed equal across all shards (uniform per-shard padding), matching the @@ -888,9 +888,12 @@ def _custom_ring_attention_forward( perm: Explicit `ppermute` permutation. Defaults to a full-axis +1 rotation. For the hybrid split, pass a perm that rotates K/V *within each ring sub-group only* (built by the caller from the U x R factorization). + save_residuals: If True, returns `(out, lse)` where `lse` is the global + logsumexp of shape `(num_q_heads, orig_q_seq_len)`. Returns: - Normalized attention output, shape `(num_q_heads, q_seq_len, head_dim_v)`. + Normalized attention output of shape `(num_q_heads, q_seq_len, head_dim_v)`, + or `(out, lse)` if `save_residuals=True`. """ axis_size = lax.axis_size(ring_axis) if bidirectional: @@ -901,6 +904,8 @@ def _custom_ring_attention_forward( ) if use_fixed_m: raise NotImplementedError("fixed-m is not yet supported on the bidirectional ring path.") + if save_residuals: + raise NotImplementedError("save_residuals is not yet supported on the bidirectional ring path.") return _custom_bidirectional_ring_forward( q, k, @@ -922,53 +927,26 @@ def _custom_ring_attention_forward( shift = partial(lax.ppermute, axis_name=ring_axis, perm=perm) exp_fn = jnp.exp2 if use_base2_exp else jnp.exp + log_fn = jnp.log2 if use_base2_exp else jnp.log num_q_heads = q.shape[0] head_dim_v = v.shape[-1] if use_fixed_m: - # Fixed-m ring: each hop gates PER (head, K-shard) against the halved - # un-smoothed bound, so a head can be fixed on one shard and online on - # another. A fixed hop returns m = the Cauchy-Schwarz upper bound (not the - # rowmax); the naive (m, l) merge below would then flush the other hop's - # partial (exp(m_other - m_bound) underflows once the overshoot exceeds - # the f32 window). Merge in LSE space instead: lse = m + log(l) is - # invariant to the kernel's m convention, so overshoot cancels exactly. - # The K-shard norms rotate WITH K/V (a (heads,)-sized ppermute) instead of - # being re-reduced per hop, which would stall the kernel's scalar prefetch. + if save_residuals: + raise NotImplementedError("save_residuals is not supported with use_fixed_m.") if fixed_m_norms is None: raise ValueError("use_fixed_m on the ring path requires fixed_m_norms=(qn_max, mk_h).") - log_fn = jnp.log2 if use_base2_exp else jnp.log qn_max, mk_h_init = fixed_m_norms tiny = jnp.finfo(jnp.float32).tiny - # Finite (not -inf) init: the first merge computes exp(init - lse_new) = 0.0 - # exactly; a -inf init meeting an empty partial would produce inf - inf = NaN. lse_init = -1e30 - # Every rank's K-shard norms, gathered ONCE before the scan: (R, heads). - # Rotating mk alongside K/V instead (a third per-hop ppermute feeding the - # kernel's scalar prefetch) serialized the K/V rotation against the kernel - # (trace: collective-permute-done 0.004s -> 0.467s per window); a local - # index into a pre-gathered array keeps the per-hop gate collective-free. - # A caller holding the full table already (e.g. a static weight-derived - # bound, identical on every rank) passes it as (ring_size, heads) and - # skips the gather -- an all_gather of a constant is NOT folded by XLA - # and would still occupy the async-collective machinery every call. if mk_h_init.ndim == 2: mk_all = mk_h_init else: mk_all = lax.all_gather(mk_h_init, ring_axis) # (axis_size, heads) my_ring_index = lax.axis_index(ring_axis) - # GLOBAL bound = max over every shard's mk. When ALL local heads pass the - # gate at this single bound, every hop's pinned m is IDENTICAL (it depends - # only on the stationary q rows and the global bound), so hop partials - # combine by PURE ACCUMULATION: o += o_hop, l += l_hop, one normalize at - # the end -- no per-hop LSE math or [H,S,D] divides. The predicate is - # device-uniform ALONG THE RING (the caller pmaxes qn over the ring axis - # and mk_all is a gathered table), so every ppermute participant takes the - # same lax.cond branch; ulysses ranks may diverge freely (no ulysses - # collective lives inside the branches). mk_global = mk_all.max(axis=0) # (heads,) all_fixed_global = jnp.all( qn_max * mk_global <= custom_splash._FIXED_M_RING_SAFE_BOUND # pylint: disable=protected-access @@ -979,18 +957,9 @@ def _accumulate_scan(_): o_sum = jnp.zeros((num_q_heads, orig_q_seq_len, head_dim_v), jnp.float32) l_sum = jnp.zeros((num_q_heads, orig_q_seq_len), jnp.float32) k_current, v_current = k, v - # Python loop over the (static) ring size rather than a scan: it lets the - # LAST hop skip its rotation. A scan body must rotate unconditionally, and - # the trailing ppermute is NOT dead-code-eliminated (collectives carry - # cross-device pairing), so a scan pays ring_size rotations to consume - # ring_size - 1 shards. At 2 x 387MB/hop over ~16 GB/s of unidirectional - # ICI that wasted hop is ~24ms/layer of wire time -- enough to make the - # ring ICI-bound and hide any kernel-side win. for hop in range(ring_size): is_last_hop = hop == ring_size - 1 if not is_last_hop: - # Issue the next shard's rotation before computing on this one so the - # transfer overlaps the kernel. k_next = shift(k_current) v_next = shift(v_current) o_curr, _, l_curr = custom_splash._splash_attention_forward_ring( # pylint: disable=protected-access @@ -1005,11 +974,6 @@ def _accumulate_scan(_): vmem_limit_bytes=vmem_limit_bytes, use_fixed_m=True, mk=mk_arr, - # This branch only runs under `all_fixed_global`, so the kernel is - # told at compile time that every head is fixed: no per-head scalar - # dispatch, and -- load-bearing -- a single body in the ragged last - # KV block instead of two (a two-body last block is what triggers - # the Mosaic scheduler cliff on the whole grid). uniform_fixed_m=True, ) o_sum = o_sum + o_curr.astype(jnp.float32) @@ -1017,25 +981,16 @@ def _accumulate_scan(_): if not is_last_hop: k_current, v_current = k_next, v_next l_inv = jnp.where(l_sum == 0.0, 0.0, 1.0 / l_sum) - # Narrow INSIDE the branch: lax.cond has to move its output between the - # branch buffer and its own, and this array is [heads, seq, head_dim] -- - # 387MB in f32. Casting here halves that copy (measured 1.04ms/layer for - # the conditional, ~half of fixed-m's whole kernel win). return (o_sum * l_inv[..., None]).astype(q.dtype) def fixed_body(carry, hop, is_last_hop): o_run, lse_run, k_current, v_current = carry - # Prefetch the next shard while computing on this one. The last hop skips - # it: nothing consumes the rotated shard, and the collective would still - # occupy ICI (see the accumulate path's note). if is_last_hop: k_next, v_next = k_current, v_current else: k_next = shift(k_current) v_next = shift(v_current) - # perm src i -> dst i+1: after `hop` shifts this rank holds the K shard - # of ring rank (my_index - hop) mod R; its norms come from the local table. mk_h = jax.lax.dynamic_index_in_dim(mk_all, (my_ring_index - hop) % axis_size, keepdims=False) fixed_ok = (qn_max * mk_h <= custom_splash._FIXED_M_RING_SAFE_BOUND).astype( # pylint: disable=protected-access jnp.float32 @@ -1059,12 +1014,10 @@ def fixed_body(carry, hop, is_last_hop): l_curr = l_curr.astype(jnp.float32) o_curr = o_curr.astype(jnp.float32) - # Partial -> (normalized output, LSE); empty rows map to lse = -inf. l_safe = jnp.maximum(l_curr, tiny) lse_curr = jnp.where(l_curr > 0.0, m_curr + log_fn(l_safe), -jnp.inf) o_norm = o_curr / l_safe[..., None] - # LSE-space merge of two normalized partials over disjoint KV sets. lse_new = jnp.maximum(lse_run, lse_curr) w_run = exp_fn(lse_run - lse_new) w_curr = exp_fn(lse_curr - lse_new) @@ -1128,9 +1081,209 @@ def _lse_scan(_): l_inv = jnp.where(l_final == 0.0, 0.0, 1.0 / l_final) out = (o_final * l_inv[..., None]).astype(q.dtype) + if save_residuals: + lse = m_final + log_fn(jnp.maximum(l_final, jnp.finfo(jnp.float32).tiny)) + lse = jnp.where(l_final == 0.0, mask_value, lse) + return out, lse return out +def _custom_ring_attention_backward( + q: jax.Array, + k: jax.Array, + v: jax.Array, + out: jax.Array, + lse: jax.Array, + do: jax.Array, + *, + block_sizes: "custom_splash._BlockSizes", + orig_q_seq_len: int, + orig_kv_seq_len: int, + use_base2_exp: bool, + use_experimental_scheduler: bool, + vmem_limit_bytes: int | None, + ring_axis: str, + ring_size: int | None = None, + perm: list[tuple[int, int]] | tuple[tuple[int, int], ...] | None = None, +) -> tuple[jax.Array, jax.Array, jax.Array]: + """Backward ring attention using the custom dense splash backward kernel.""" + axis_size = lax.axis_size(ring_axis) + if ring_size is None: + ring_size = axis_size + if perm is None: + perm = [(i, (i + 1) % axis_size) for i in range(axis_size)] + + shift = partial(lax.ppermute, axis_name=ring_axis, perm=perm) + + di_vec = jnp.sum(out.astype(jnp.float32) * do.astype(jnp.float32), axis=-1) + di_expanded = jnp.broadcast_to( + di_vec[:, None, :], (q.shape[0], custom_splash.NUM_SUBLANES, orig_q_seq_len) + ) + out_swapped = jnp.swapaxes(out, 1, 2) + do_swapped = jnp.swapaxes(do, 1, 2) + lse_expanded = jnp.broadcast_to( + lse[:, None, :], (q.shape[0], custom_splash.NUM_SUBLANES, orig_q_seq_len) + ) + + dq_accum = jnp.zeros_like(q, dtype=jnp.float32) + dk_accum = jnp.zeros_like(k, dtype=jnp.float32) + dv_accum = jnp.zeros_like(v, dtype=jnp.float32) + + k_current, v_current = k, v + for hop in range(ring_size): + if hop != ring_size - 1: + k_next = shift(k_current) + v_next = shift(v_current) + + dq_i, dk_i, dv_i = custom_splash._splash_attention_backward( # pylint: disable=protected-access + q, + k_current, + v_current, + out_swapped, + lse_expanded, + do_swapped, + block_sizes, + q_seq_len=orig_q_seq_len, + kv_seq_len=orig_kv_seq_len, + use_base2_exp=use_base2_exp, + use_experimental_scheduler=use_experimental_scheduler, + vmem_limit_bytes=vmem_limit_bytes, + di=di_expanded, + ) + dq_accum = dq_accum + dq_i.astype(jnp.float32) + dk_accum = dk_accum + dk_i.astype(jnp.float32) + dv_accum = dv_accum + dv_i.astype(jnp.float32) + + if ring_size > 1: + dk_accum = shift(dk_accum) + dv_accum = shift(dv_accum) + if hop != ring_size - 1: + k_current, v_current = k_next, v_next + + return dq_accum.astype(q.dtype), dk_accum.astype(k.dtype), dv_accum.astype(v.dtype) + + +@partial( + jax.custom_vjp, + nondiff_argnames=( + "block_sizes", + "orig_q_seq_len", + "orig_kv_seq_len", + "use_base2_exp", + "use_experimental_scheduler", + "vmem_limit_bytes", + "mask_value", + "ring_axis", + "ring_size", + "perm", + ), +) +def _custom_ring_attention_custom( + q: jax.Array, + k: jax.Array, + v: jax.Array, + block_sizes: "custom_splash._BlockSizes", + orig_q_seq_len: int, + orig_kv_seq_len: int, + use_base2_exp: bool, + use_experimental_scheduler: bool, + vmem_limit_bytes: int | None, + mask_value: float, + ring_axis: str, + ring_size: int | None, + perm: tuple[tuple[int, int], ...] | None, +) -> jax.Array: + return _custom_ring_attention_forward( + q, + k, + v, + block_sizes=block_sizes, + orig_q_seq_len=orig_q_seq_len, + orig_kv_seq_len=orig_kv_seq_len, + use_base2_exp=use_base2_exp, + use_experimental_scheduler=use_experimental_scheduler, + vmem_limit_bytes=vmem_limit_bytes, + mask_value=mask_value, + ring_axis=ring_axis, + ring_size=ring_size, + perm=list(perm) if perm is not None else None, + save_residuals=False, + ) + + +def _custom_ring_attention_fwd( + q: jax.Array, + k: jax.Array, + v: jax.Array, + block_sizes: "custom_splash._BlockSizes", + orig_q_seq_len: int, + orig_kv_seq_len: int, + use_base2_exp: bool, + use_experimental_scheduler: bool, + vmem_limit_bytes: int | None, + mask_value: float, + ring_axis: str, + ring_size: int | None, + perm: tuple[tuple[int, int], ...] | None, +): + out, lse = _custom_ring_attention_forward( + q, + k, + v, + block_sizes=block_sizes, + orig_q_seq_len=orig_q_seq_len, + orig_kv_seq_len=orig_kv_seq_len, + use_base2_exp=use_base2_exp, + use_experimental_scheduler=use_experimental_scheduler, + vmem_limit_bytes=vmem_limit_bytes, + mask_value=mask_value, + ring_axis=ring_axis, + ring_size=ring_size, + perm=list(perm) if perm is not None else None, + save_residuals=True, + ) + return out, (q, k, v, out, lse) + + +def _custom_ring_attention_bwd( + block_sizes: "custom_splash._BlockSizes", + orig_q_seq_len: int, + orig_kv_seq_len: int, + use_base2_exp: bool, + use_experimental_scheduler: bool, + vmem_limit_bytes: int | None, + mask_value: float, + ring_axis: str, + ring_size: int | None, + perm: tuple[tuple[int, int], ...] | None, + residuals, + do: jax.Array, +): + del mask_value + q, k, v, out, lse = residuals + dq, dk, dv = _custom_ring_attention_backward( + q, + k, + v, + out, + lse, + do, + block_sizes=block_sizes, + orig_q_seq_len=orig_q_seq_len, + orig_kv_seq_len=orig_kv_seq_len, + use_base2_exp=use_base2_exp, + use_experimental_scheduler=use_experimental_scheduler, + vmem_limit_bytes=vmem_limit_bytes, + ring_axis=ring_axis, + ring_size=ring_size, + perm=list(perm) if perm is not None else None, + ) + return dq, dk, dv + + +_custom_ring_attention_custom.defvjp(_custom_ring_attention_fwd, _custom_ring_attention_bwd) + + def make_custom_ring_attention( *, block_sizes: "custom_splash._BlockSizes", @@ -1147,7 +1300,7 @@ def make_custom_ring_attention( use_fixed_m: bool = False, fixed_m_norms: tuple[jax.Array, jax.Array] | None = None, ): - """Builds a forward-only ring-attention callable around the custom kernel. + """Builds a ring-attention callable around the custom kernel (supports forward & backward). The returned function takes a single (un-batched) `(q, k, v)` triple of shape `(num_heads, seq, head_dim)` and is meant to be `jax.vmap`-ped over the batch @@ -1162,9 +1315,29 @@ def make_custom_ring_attention( one hop at a time) for a NON-wrapping ring axis, avoiding the diameter-length wrap hop. Requires `perm=None` and the full real ring axis (no sub-group). """ + perm_tuple = tuple(tuple(x) for x in perm) if perm is not None else None def _ring(q, k, v): - return _custom_ring_attention_forward( + if use_fixed_m or bidirectional: + return _custom_ring_attention_forward( + q, + k, + v, + block_sizes=block_sizes, + orig_q_seq_len=orig_q_seq_len, + orig_kv_seq_len=orig_kv_seq_len, + use_base2_exp=use_base2_exp, + use_experimental_scheduler=use_experimental_scheduler, + vmem_limit_bytes=vmem_limit_bytes, + mask_value=mask_value, + ring_axis=ring_axis, + ring_size=ring_size, + perm=perm, + bidirectional=bidirectional, + use_fixed_m=use_fixed_m, + fixed_m_norms=fixed_m_norms, + ) + return _custom_ring_attention_custom( q, k, v, @@ -1177,10 +1350,7 @@ def _ring(q, k, v): mask_value=mask_value, ring_axis=ring_axis, ring_size=ring_size, - perm=perm, - bidirectional=bidirectional, - use_fixed_m=use_fixed_m, - fixed_m_norms=fixed_m_norms, + perm=perm_tuple, ) return _ring diff --git a/src/maxdiffusion/tests/custom_splash_backward_test.py b/src/maxdiffusion/tests/custom_splash_backward_test.py new file mode 100644 index 000000000..6a47cf244 --- /dev/null +++ b/src/maxdiffusion/tests/custom_splash_backward_test.py @@ -0,0 +1,471 @@ +""" +Copyright 2026 Google LLC + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + https://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +""" + +"""Unit tests for custom splash attention and custom ring attention backward pass.""" + +import functools +import math +import unittest + +import jax +import jax.numpy as jnp +import numpy as np + +from maxdiffusion.kernels import custom_splash_attention as custom_splash +from maxdiffusion.kernels.splash_attention import ring_attention_kernel + +_LOG2E = math.log2(math.e) +_LN2 = math.log(2.0) + + +def _reference_attention(q, k, v, actual_q_len, actual_kv_len, use_base2_exp): + """Reference attention in fp32 returning shape (Hq, Dv, actual_q_len).""" + q_f = q[:, :actual_q_len, :].astype(jnp.float32) + k_f = k[:, :actual_kv_len, :].astype(jnp.float32) + v_f = v[:, :actual_kv_len, :].astype(jnp.float32) + hq = q_f.shape[0] + hkv = k_f.shape[0] + q_per_kv = hq // hkv + k_f = jnp.repeat(k_f, q_per_kv, axis=0) + v_f = jnp.repeat(v_f, q_per_kv, axis=0) + logits = jnp.einsum("hsd,htd->hst", q_f, k_f) + if use_base2_exp: + logits = logits * _LN2 + probs = jax.nn.softmax(logits, axis=-1) + out = jnp.einsum("hst,htd->hds", probs, v_f) + return out.astype(q.dtype) + + +class CustomSplashBackwardTest(unittest.TestCase): + """Tests custom splash attention backward pass and ring attention backward pass.""" + + def setUp(self): + super().setUp() + if jax.default_backend() != "tpu": + self.skipTest("Custom Pallas splash kernel requires TPU.") + self.block_sizes = custom_splash._BlockSizes( + block_q=1024, + block_kv=1024, + block_kv_compute=512, + block_kv_compute_in=256, + ) + + def test_single_device_bwd_exact_multiples(self): + hq, hkv, sq, skv, d = 4, 4, 2048, 2048, 128 + scale = 1.0 / math.sqrt(d) + k1, k2, k3, k4 = jax.random.split(jax.random.PRNGKey(0), 4) + q = (jax.random.normal(k1, (hq, sq, d), jnp.bfloat16) * scale * _LOG2E).astype(jnp.bfloat16) + k = jax.random.normal(k2, (hkv, skv, d), jnp.bfloat16) * scale + v = jax.random.normal(k3, (hkv, skv, d), jnp.bfloat16) * scale + do = jax.random.normal(k4, (hq, d, sq), jnp.bfloat16) + + kernel = custom_splash.make_splash_mha( + block_sizes=self.block_sizes, + orig_q_seq_len=sq, + orig_kv_seq_len=skv, + use_base2_exp=True, + ) + + out_kern, vjp_kern = jax.vjp(kernel, q, k, v) + dq_kern, dk_kern, dv_kern = vjp_kern(do) + + out_ref, vjp_ref = jax.vjp( + lambda q_, k_, v_: _reference_attention(q_, k_, v_, sq, skv, True), + q, + k, + v, + ) + dq_ref, dk_ref, dv_ref = vjp_ref(do) + + self.assertLess(float(jnp.max(jnp.abs(out_kern.astype(jnp.float32) - out_ref.astype(jnp.float32)))), 1e-3) + self.assertLess(float(jnp.max(jnp.abs(dq_kern.astype(jnp.float32) - dq_ref.astype(jnp.float32)))), 1e-3) + self.assertLess(float(jnp.max(jnp.abs(dk_kern.astype(jnp.float32) - dk_ref.astype(jnp.float32)))), 1e-3) + self.assertLess(float(jnp.max(jnp.abs(dv_kern.astype(jnp.float32) - dv_ref.astype(jnp.float32)))), 2e-3) + + def test_single_device_bwd_ragged_gqa_and_vmapped(self): + batch, hq, hkv, sq, skv, d = 2, 4, 2, 2048, 2048, 128 + actual_q_len, actual_kv_len = 1500, 1600 + scale = 1.0 / math.sqrt(d) + k1, k2, k3, k4 = jax.random.split(jax.random.PRNGKey(1), 4) + q = (jax.random.normal(k1, (batch, hq, sq, d), jnp.bfloat16) * scale * _LOG2E).astype(jnp.bfloat16) + k = jax.random.normal(k2, (batch, hkv, skv, d), jnp.bfloat16) * scale + v = jax.random.normal(k3, (batch, hkv, skv, d), jnp.bfloat16) * scale + do = jax.random.normal(k4, (batch, hq, d, actual_q_len), jnp.bfloat16) + + kernel = custom_splash.make_splash_mha( + block_sizes=self.block_sizes, + orig_q_seq_len=actual_q_len, + orig_kv_seq_len=actual_kv_len, + use_base2_exp=True, + ) + vmapped_kernel = jax.vmap(kernel, in_axes=(0, 0, 0)) + + out_kern, vjp_kern = jax.vjp(vmapped_kernel, q, k, v) + dq_kern, dk_kern, dv_kern = vjp_kern(do) + + vmapped_ref = jax.vmap( + lambda q_, k_, v_: _reference_attention(q_, k_, v_, actual_q_len, actual_kv_len, True), + in_axes=(0, 0, 0), + ) + out_ref, vjp_ref = jax.vjp(vmapped_ref, q, k, v) + dq_ref, dk_ref, dv_ref = vjp_ref(do) + + self.assertLess(float(jnp.max(jnp.abs(out_kern.astype(jnp.float32) - out_ref.astype(jnp.float32)))), 1e-3) + self.assertLess( + float(jnp.max(jnp.abs(dq_kern[:, :, :actual_q_len].astype(jnp.float32) - dq_ref[:, :, :actual_q_len].astype(jnp.float32)))), + 1e-3, + ) + self.assertEqual(float(jnp.max(jnp.abs(dq_kern[:, :, actual_q_len:].astype(jnp.float32)))), 0.0) + self.assertLess( + float(jnp.max(jnp.abs(dk_kern[:, :, :actual_kv_len].astype(jnp.float32) - dk_ref[:, :, :actual_kv_len].astype(jnp.float32)))), + 1e-3, + ) + self.assertEqual(float(jnp.max(jnp.abs(dk_kern[:, :, actual_kv_len:].astype(jnp.float32)))), 0.0) + self.assertLess( + float(jnp.max(jnp.abs(dv_kern[:, :, :actual_kv_len].astype(jnp.float32) - dv_ref[:, :, :actual_kv_len].astype(jnp.float32)))), + 2e-3, + ) + self.assertEqual(float(jnp.max(jnp.abs(dv_kern[:, :, actual_kv_len:].astype(jnp.float32)))), 0.0) + + def test_single_device_bwd_fused_aliasing_and_unfused(self): + hq, hkv, sq, skv, d = 4, 2, 2048, 4096, 128 + actual_q_len, actual_kv_len = 1500, 3500 # grid_width = ceil(3500/512) = 7 > 3 + scale = 1.0 / math.sqrt(d) + k1, k2, k3, k4 = jax.random.split(jax.random.PRNGKey(42), 4) + q = (jax.random.normal(k1, (hq, sq, d), jnp.bfloat16) * scale * _LOG2E).astype(jnp.bfloat16) + k = jax.random.normal(k2, (hkv, skv, d), jnp.bfloat16) * scale + v = jax.random.normal(k3, (hkv, skv, d), jnp.bfloat16) * scale + do = jax.random.normal(k4, (hq, d, actual_q_len), jnp.bfloat16) + + out_ref, vjp_ref = jax.vjp( + lambda q_, k_, v_: _reference_attention(q_, k_, v_, actual_q_len, actual_kv_len, True), + q, + k, + v, + ) + dq_ref, dk_ref, dv_ref = vjp_ref(do) + + for use_fused, dq_red in [(True, 3), (True, None), (False, None)]: + bs = custom_splash._BlockSizes( + block_q=512, + block_kv=512, + block_kv_compute=256, + block_kv_compute_in=256, + use_fused_bwd_kernel=use_fused, + dq_reduction_steps=dq_red, + ) + kernel = custom_splash.make_splash_mha( + block_sizes=bs, + orig_q_seq_len=actual_q_len, + orig_kv_seq_len=actual_kv_len, + use_base2_exp=True, + ) + out_k, vjp_k = jax.vjp(kernel, q, k, v) + dq_k, dk_k, dv_k = vjp_k(do) + self.assertLess(float(jnp.max(jnp.abs(out_k.astype(jnp.float32) - out_ref.astype(jnp.float32)))), 1e-3) + self.assertLess( + float(jnp.max(jnp.abs(dq_k[:, :actual_q_len].astype(jnp.float32) - dq_ref[:, :actual_q_len].astype(jnp.float32)))), + 1e-3, + ) + self.assertEqual(float(jnp.max(jnp.abs(dq_k[:, actual_q_len:].astype(jnp.float32)))), 0.0) + self.assertLess( + float(jnp.max(jnp.abs(dk_k[:, :actual_kv_len].astype(jnp.float32) - dk_ref[:, :actual_kv_len].astype(jnp.float32)))), + 1e-3, + ) + self.assertEqual(float(jnp.max(jnp.abs(dk_k[:, actual_kv_len:].astype(jnp.float32)))), 0.0) + self.assertLess( + float(jnp.max(jnp.abs(dv_k[:, :actual_kv_len].astype(jnp.float32) - dv_ref[:, :actual_kv_len].astype(jnp.float32)))), + 2e-3, + ) + self.assertEqual(float(jnp.max(jnp.abs(dv_k[:, actual_kv_len:].astype(jnp.float32)))), 0.0) + + def test_ring_attention_bwd_ragged_gqa_vmapped(self): + devices = jax.devices() + if len(devices) < 4: + self.skipTest("Requires at least 4 TPU devices for ring attention test.") + ring_size = 4 + mesh = jax.sharding.Mesh(np.array(devices[:ring_size]), ("context",)) + batch, hq, hkv, sq_local, skv_local, d = 2, 4, 2, 1024, 1024, 128 + actual_q_local, actual_kv_local = 800, 900 + bsizes = custom_splash._BlockSizes( + block_q=512, + block_kv=512, + block_kv_compute=256, + block_kv_compute_in=256, + ) + + scale = 1.0 / math.sqrt(d) + k1, k2, k3, k4 = jax.random.split(jax.random.PRNGKey(2), 4) + q_global = (jax.random.normal(k1, (ring_size, batch, hq, sq_local, d), jnp.bfloat16) * scale * _LOG2E).astype(jnp.bfloat16) + k_global = jax.random.normal(k2, (ring_size, batch, hkv, skv_local, d), jnp.bfloat16) * scale + v_global = jax.random.normal(k3, (ring_size, batch, hkv, skv_local, d), jnp.bfloat16) * scale + do_global = jax.random.normal(k4, (ring_size, batch, hq, actual_q_local, d), jnp.bfloat16) + + def _ref_single_batch(qg, kg, vg): + q_list = [qg[r, :, :actual_q_local, :].astype(jnp.float32) for r in range(ring_size)] + k_list = [kg[r, :, :actual_kv_local, :].astype(jnp.float32) for r in range(ring_size)] + v_list = [vg[r, :, :actual_kv_local, :].astype(jnp.float32) for r in range(ring_size)] + q_cat = jnp.concatenate(q_list, axis=1) + k_cat = jnp.concatenate(k_list, axis=1) + v_cat = jnp.concatenate(v_list, axis=1) + k_cat = jnp.repeat(k_cat, hq // hkv, axis=0) + v_cat = jnp.repeat(v_cat, hq // hkv, axis=0) + logits = jnp.einsum("hsd,htd->hst", q_cat, k_cat) * _LN2 + probs = jax.nn.softmax(logits, axis=-1) + out_cat = jnp.einsum("hst,htd->hsd", probs, v_cat) + return jnp.stack( + [out_cat[:, r * actual_q_local : (r + 1) * actual_q_local, :] for r in range(ring_size)], + axis=0, + ).astype(qg.dtype) + + ref_fn = jax.vmap(_ref_single_batch, in_axes=(1, 1, 1), out_axes=1) + out_ref, vjp_ref = jax.vjp(ref_fn, q_global, k_global, v_global) + dq_ref, dk_ref, dv_ref = vjp_ref(do_global) + + for use_fused in [True, False]: + bsizes = custom_splash._BlockSizes( + block_q=512, + block_kv=512, + block_kv_compute=256, + block_kv_compute_in=256, + use_fused_bwd_kernel=use_fused, + ) + ring_kernel = ring_attention_kernel.make_custom_ring_attention( + block_sizes=bsizes, + orig_q_seq_len=actual_q_local, + orig_kv_seq_len=actual_kv_local, + use_base2_exp=True, + ring_axis="context", + ) + vmapped_ring = jax.vmap(ring_kernel, in_axes=(0, 0, 0)) + + p = jax.sharding.PartitionSpec("context", None, None, None, None) + + @functools.partial( + jax.shard_map, + mesh=mesh, + in_specs=(p, p, p, p), + out_specs=(p, p, p, p), + check_vma=False, + ) + def run_ring(q_sh, k_sh, v_sh, do_sh): + q_s, k_s, v_s, do_s = q_sh[0], k_sh[0], v_sh[0], do_sh[0] + out_s, vjp_s = jax.vjp(vmapped_ring, q_s, k_s, v_s) + dq_s, dk_s, dv_s = vjp_s(do_s) + return out_s[None], dq_s[None], dk_s[None], dv_s[None] + + out_ring, dq_ring, dk_ring, dv_ring = run_ring(q_global, k_global, v_global, do_global) + + self.assertLess(float(jnp.max(jnp.abs(out_ring.astype(jnp.float32) - out_ref.astype(jnp.float32)))), 1e-3) + self.assertLess( + float(jnp.max(jnp.abs(dq_ring[:, :, :, :actual_q_local].astype(jnp.float32) - dq_ref[:, :, :, :actual_q_local].astype(jnp.float32)))), + 1e-3, + ) + self.assertEqual(float(jnp.max(jnp.abs(dq_ring[:, :, :, actual_q_local:].astype(jnp.float32)))), 0.0) + self.assertLess( + float(jnp.max(jnp.abs(dk_ring[:, :, :, :actual_kv_local].astype(jnp.float32) - dk_ref[:, :, :, :actual_kv_local].astype(jnp.float32)))), + 1e-3, + ) + self.assertEqual(float(jnp.max(jnp.abs(dk_ring[:, :, :, actual_kv_local:].astype(jnp.float32)))), 0.0) + self.assertLess( + float(jnp.max(jnp.abs(dv_ring[:, :, :, :actual_kv_local].astype(jnp.float32) - dv_ref[:, :, :, :actual_kv_local].astype(jnp.float32)))), + 2e-3, + ) + self.assertEqual(float(jnp.max(jnp.abs(dv_ring[:, :, :, actual_kv_local:].astype(jnp.float32)))), 0.0) + + def test_single_device_bwd_unpadded_non_divisible_seqlens(self): + """Tests directly passing unpadded Q/K/V whose sequence lengths are not divisible by block size.""" + hq, hkv, sq, skv, d = 4, 2, 1350, 1777, 128 + scale = 1.0 / math.sqrt(d) + k1, k2, k3, k4 = jax.random.split(jax.random.PRNGKey(101), 4) + q = (jax.random.normal(k1, (hq, sq, d), jnp.bfloat16) * scale * _LOG2E).astype(jnp.bfloat16) + k = jax.random.normal(k2, (hkv, skv, d), jnp.bfloat16) * scale + v = jax.random.normal(k3, (hkv, skv, d), jnp.bfloat16) * scale + do = jax.random.normal(k4, (hq, d, sq), jnp.bfloat16) + + out_ref, vjp_ref = jax.vjp( + lambda q_, k_, v_: _reference_attention(q_, k_, v_, sq, skv, True), + q, + k, + v, + ) + dq_ref, dk_ref, dv_ref = vjp_ref(do) + + for use_fused, dq_red in [(True, 3), (True, None), (False, None)]: + bs = custom_splash._BlockSizes( + block_q=512, + block_kv=512, + block_kv_compute=256, + block_kv_compute_in=128, + use_fused_bwd_kernel=use_fused, + dq_reduction_steps=dq_red, + ) + kernel = custom_splash.make_splash_mha( + block_sizes=bs, + orig_q_seq_len=sq, + orig_kv_seq_len=skv, + use_base2_exp=True, + ) + out_k, vjp_k = jax.vjp(kernel, q, k, v) + dq_k, dk_k, dv_k = vjp_k(do) + + self.assertEqual(out_k.shape, (hq, d, sq)) + self.assertEqual(dq_k.shape, (hq, sq, d)) + self.assertEqual(dk_k.shape, (hkv, skv, d)) + self.assertEqual(dv_k.shape, (hkv, skv, d)) + + self.assertLess(float(jnp.max(jnp.abs(out_k.astype(jnp.float32) - out_ref.astype(jnp.float32)))), 1e-3) + self.assertLess(float(jnp.max(jnp.abs(dq_k.astype(jnp.float32) - dq_ref.astype(jnp.float32)))), 1e-3) + self.assertLess(float(jnp.max(jnp.abs(dk_k.astype(jnp.float32) - dk_ref.astype(jnp.float32)))), 1e-3) + self.assertLess(float(jnp.max(jnp.abs(dv_k.astype(jnp.float32) - dv_ref.astype(jnp.float32)))), 2e-3) + + def test_single_device_bwd_padded_with_garbage_in_kv_tail(self): + """Verifies that non-divisible orig_kv_seq_len ignores extreme garbage values in padded KV tail.""" + hq, hkv, padded_sq, padded_skv, d = 4, 2, 2048, 2048, 128 + actual_q_len, actual_kv_len = 1350, 1777 + scale = 1.0 / math.sqrt(d) + k1, k2, k3, k4 = jax.random.split(jax.random.PRNGKey(102), 4) + q = (jax.random.normal(k1, (hq, padded_sq, d), jnp.bfloat16) * scale * _LOG2E).astype(jnp.bfloat16) + k = jax.random.normal(k2, (hkv, padded_skv, d), jnp.bfloat16) * scale + v = jax.random.normal(k3, (hkv, padded_skv, d), jnp.bfloat16) * scale + # Inject huge garbage numbers into the KV padding tail [actual_kv_len:] + k = k.at[:, actual_kv_len:, :].set(1000.0) + v = v.at[:, actual_kv_len:, :].set(1000.0) + do = jax.random.normal(k4, (hq, d, actual_q_len), jnp.bfloat16) + + out_ref, vjp_ref = jax.vjp( + lambda q_, k_, v_: _reference_attention(q_, k_, v_, actual_q_len, actual_kv_len, True), + q, + k, + v, + ) + dq_ref, dk_ref, dv_ref = vjp_ref(do) + + for use_fused in [True, False]: + bs = custom_splash._BlockSizes( + block_q=512, + block_kv=512, + block_kv_compute=256, + block_kv_compute_in=128, + use_fused_bwd_kernel=use_fused, + ) + kernel = custom_splash.make_splash_mha( + block_sizes=bs, + orig_q_seq_len=actual_q_len, + orig_kv_seq_len=actual_kv_len, + use_base2_exp=True, + ) + out_k, vjp_k = jax.vjp(kernel, q, k, v) + dq_k, dk_k, dv_k = vjp_k(do) + + self.assertLess(float(jnp.max(jnp.abs(out_k.astype(jnp.float32) - out_ref.astype(jnp.float32)))), 1e-3) + self.assertLess( + float(jnp.max(jnp.abs(dq_k[:, :actual_q_len].astype(jnp.float32) - dq_ref[:, :actual_q_len].astype(jnp.float32)))), + 1e-3, + ) + self.assertEqual(float(jnp.max(jnp.abs(dq_k[:, actual_q_len:].astype(jnp.float32)))), 0.0) + self.assertLess( + float(jnp.max(jnp.abs(dk_k[:, :actual_kv_len].astype(jnp.float32) - dk_ref[:, :actual_kv_len].astype(jnp.float32)))), + 1e-3, + ) + self.assertEqual(float(jnp.max(jnp.abs(dk_k[:, actual_kv_len:].astype(jnp.float32)))), 0.0) + self.assertLess( + float(jnp.max(jnp.abs(dv_k[:, :actual_kv_len].astype(jnp.float32) - dv_ref[:, :actual_kv_len].astype(jnp.float32)))), + 2e-3, + ) + self.assertEqual(float(jnp.max(jnp.abs(dv_k[:, actual_kv_len:].astype(jnp.float32)))), 0.0) + + def test_ring_attention_bwd_unpadded_non_divisible_seqlens(self): + """Tests 4-device ring attention with unpadded local sequence lengths not divisible by block sizes.""" + devices = jax.devices() + if len(devices) < 4: + self.skipTest("Requires at least 4 TPU devices for ring attention test.") + ring_size = 4 + mesh = jax.sharding.Mesh(np.array(devices[:ring_size]), ("context",)) + batch, hq, hkv, sq_local, skv_local, d = 2, 4, 2, 650, 733, 128 + + scale = 1.0 / math.sqrt(d) + k1, k2, k3, k4 = jax.random.split(jax.random.PRNGKey(103), 4) + q_global = (jax.random.normal(k1, (ring_size, batch, hq, sq_local, d), jnp.bfloat16) * scale * _LOG2E).astype(jnp.bfloat16) + k_global = jax.random.normal(k2, (ring_size, batch, hkv, skv_local, d), jnp.bfloat16) * scale + v_global = jax.random.normal(k3, (ring_size, batch, hkv, skv_local, d), jnp.bfloat16) * scale + do_global = jax.random.normal(k4, (ring_size, batch, hq, sq_local, d), jnp.bfloat16) + + def _ref_single_batch(qg, kg, vg): + q_cat = jnp.concatenate([qg[r].astype(jnp.float32) for r in range(ring_size)], axis=1) + k_cat = jnp.concatenate([kg[r].astype(jnp.float32) for r in range(ring_size)], axis=1) + v_cat = jnp.concatenate([vg[r].astype(jnp.float32) for r in range(ring_size)], axis=1) + k_cat = jnp.repeat(k_cat, hq // hkv, axis=0) + v_cat = jnp.repeat(v_cat, hq // hkv, axis=0) + logits = jnp.einsum("hsd,htd->hst", q_cat, k_cat) * _LN2 + probs = jax.nn.softmax(logits, axis=-1) + out_cat = jnp.einsum("hst,htd->hsd", probs, v_cat) + return jnp.stack( + [out_cat[:, r * sq_local : (r + 1) * sq_local, :] for r in range(ring_size)], + axis=0, + ).astype(qg.dtype) + + ref_fn = jax.vmap(_ref_single_batch, in_axes=(1, 1, 1), out_axes=1) + out_ref, vjp_ref = jax.vjp(ref_fn, q_global, k_global, v_global) + dq_ref, dk_ref, dv_ref = vjp_ref(do_global) + + for use_fused in [True, False]: + bsizes = custom_splash._BlockSizes( + block_q=512, + block_kv=512, + block_kv_compute=256, + block_kv_compute_in=128, + use_fused_bwd_kernel=use_fused, + ) + ring_kernel = ring_attention_kernel.make_custom_ring_attention( + block_sizes=bsizes, + orig_q_seq_len=sq_local, + orig_kv_seq_len=skv_local, + use_base2_exp=True, + ring_axis="context", + ) + vmapped_ring = jax.vmap(ring_kernel, in_axes=(0, 0, 0)) + + p = jax.sharding.PartitionSpec("context", None, None, None, None) + + @functools.partial( + jax.shard_map, + mesh=mesh, + in_specs=(p, p, p, p), + out_specs=(p, p, p, p), + check_vma=False, + ) + def run_ring(q_sh, k_sh, v_sh, do_sh): + q_s, k_s, v_s, do_s = q_sh[0], k_sh[0], v_sh[0], do_sh[0] + out_s, vjp_s = jax.vjp(vmapped_ring, q_s, k_s, v_s) + dq_s, dk_s, dv_s = vjp_s(do_s) + return out_s[None], dq_s[None], dk_s[None], dv_s[None] + + out_ring, dq_ring, dk_ring, dv_ring = run_ring(q_global, k_global, v_global, do_global) + + self.assertEqual(out_ring.shape, (ring_size, batch, hq, sq_local, d)) + self.assertEqual(dq_ring.shape, (ring_size, batch, hq, sq_local, d)) + self.assertEqual(dk_ring.shape, (ring_size, batch, hkv, skv_local, d)) + self.assertEqual(dv_ring.shape, (ring_size, batch, hkv, skv_local, d)) + + self.assertLess(float(jnp.max(jnp.abs(out_ring.astype(jnp.float32) - out_ref.astype(jnp.float32)))), 1e-3) + self.assertLess(float(jnp.max(jnp.abs(dq_ring.astype(jnp.float32) - dq_ref.astype(jnp.float32)))), 1e-3) + self.assertLess(float(jnp.max(jnp.abs(dk_ring.astype(jnp.float32) - dk_ref.astype(jnp.float32)))), 1e-3) + self.assertLess(float(jnp.max(jnp.abs(dv_ring.astype(jnp.float32) - dv_ref.astype(jnp.float32)))), 2e-3) + + +if __name__ == "__main__": + unittest.main()