From 7499339e1d749eb5caa70f700cfdc2db8ab88d4a Mon Sep 17 00:00:00 2001 From: maxtext authors Date: Tue, 21 Jul 2026 10:44:34 -0700 Subject: [PATCH] Updates to JAX flash attention kernel. - Sink jnp.dynamic-slice instructions fetching blocks of q, k and v into the inner loop. - Forced final @V layout for better utilization - Forcing fusions with must_fuse_call - Adding support to unroll the loops PiperOrigin-RevId: 951577146 --- .../kernels/attention/jax_flash_attention.py | 162 +++++++++++++----- 1 file changed, 122 insertions(+), 40 deletions(-) diff --git a/src/maxtext/kernels/attention/jax_flash_attention.py b/src/maxtext/kernels/attention/jax_flash_attention.py index 93a1933615..7308f51783 100644 --- a/src/maxtext/kernels/attention/jax_flash_attention.py +++ b/src/maxtext/kernels/attention/jax_flash_attention.py @@ -16,9 +16,14 @@ from typing import Optional, Tuple, Union import jax +from jax.experimental import layout +from jax.experimental.xla_metadata import must_fuse_call import jax.numpy as jnp from maxtext.kernels.attention import splash_attention_kernel +DLL = layout.Layout +Layout = layout.Format + SegmentIds = splash_attention_kernel.SegmentIds @@ -39,6 +44,7 @@ def flash_attention_block_masked( mask_value: float, cap: Optional[float] = None, save_residuals: bool = False, + unroll: bool = False, ) -> Union[jnp.ndarray, Tuple[jnp.ndarray, Tuple[jnp.ndarray, jnp.ndarray]]]: """Computes masked flash attention using block-sparse masking. @@ -65,6 +71,8 @@ def flash_attention_block_masked( (output, dict=(logsumexp, max_logits)). Both `logsumexp` and `max_logits` are of shape (batch_size, num_kv_heads, num_q_heads // num_kv_heads, q_seq_len). + unroll: Whether to unroll the loops. This is useful for small sequence + lengths and helps with iteration skipping for causal masks. Returns: If save_residuals is True, returns a tuple containing: @@ -99,7 +107,16 @@ def flash_attention_block_masked( segment_ids_q = segment_ids.q[:, :, None] segment_ids_kv = segment_ids.kv[:, None, :] mask_full = jnp.logical_and(mask_full, segment_ids_q == segment_ids_kv) - mask_blocked = jax.jit(mask_blocker, static_argnums=[1, 2])(mask_full, block_q, block_kv) + + # In the case of a causal mask, the compute_attention_block should be executed + # if the current block (i, j) falls within the lower triangle. This means that + # the maximum query index in block i must be greater than or equal to the + # minimum key/value index in block j. + # Max q_idx in block i: (i + 1) * block_q - 1 + # Min kv_idx in block j: j * block_kv + # Condition: (i + 1) * block_q - 1 >= j * block_kv + # Which simplifies to: (i + 1) * block_q > j * block_kv + should_compute_block = lambda i, j: (i + 1) * block_q > j * block_kv # Initialize `l` (logsumexp) and `m` (max_logits) for the online softmax. # `l` is initialized to 0 since no blocks have been processed yet and the sum @@ -127,27 +144,22 @@ def flash_attention_block_masked( # Outer loop over the key/value blocks. def outer_loop_body(j, carried): output, l, m = carried - k_j_slice = jax.lax.dynamic_slice_in_dim(k, j * block_kv, block_kv, axis=-2) - v_j_slice = jax.lax.dynamic_slice_in_dim(v, j * block_kv, block_kv, axis=-2) # Inner loop over the query blocks. def inner_loop_body(i, carried_inner): output, l, m = carried_inner - # let's get the slice of Q in N dimension - q_slice = jax.lax.dynamic_slice_in_dim(q, i * block_q, block_q, axis=-2) - # Calculates the attention computation (Q@K.T)@V with online softmax for # the current query and key/value blocks. def compute_attention_block(output, l, m): - output_i_slice = jax.lax.dynamic_slice_in_dim(output, i * block_q, block_q, axis=-2) - l_i_slice = jax.lax.dynamic_slice_in_dim(l, i * block_q, block_q, axis=-1) - m_i_slice = jax.lax.dynamic_slice_in_dim(m, i * block_q, block_q, axis=-1) - s_i_j = jnp.einsum( - "bxhqc,bxkc->bxhqk", - q_slice, - k_j_slice, - preferred_element_type=data_type, + output_i_slice = jax.lax.dynamic_slice_in_dim( + output, i * block_q, block_q, axis=-2 + ) + l_i_slice = jax.lax.dynamic_slice_in_dim( + l, i * block_q, block_q, axis=-1 + ) + m_i_slice = jax.lax.dynamic_slice_in_dim( + m, i * block_q, block_q, axis=-1 ) full_mask_i_j_slice = jax.lax.dynamic_slice( mask_full, @@ -159,12 +171,53 @@ def compute_attention_block(output, l, m): (batch_size, num_kv_heads, q_groups, block_q, block_kv), ) + k_j_slice = jax.lax.dynamic_slice_in_dim( + k, j * block_kv, block_kv, axis=-2 + ) + v_j_slice = jax.lax.dynamic_slice_in_dim( + v, j * block_kv, block_kv, axis=-2 + ) + + # let's get the slice of Q in N dimension + q_slice = jax.lax.dynamic_slice_in_dim(q, i * block_q, block_q, axis=-2) + + s_i_j_dup = jnp.einsum( + "bxhqc,bxkc->bxhqk", + q_slice, + k_j_slice, + preferred_element_type=data_type, + ) + if unroll: + if i == j: + s_i_j_dup = jnp.where(broadcasted_mask, s_i_j_dup, mask_value) + else: + s_i_j_dup = jnp.where(broadcasted_mask, s_i_j_dup, mask_value) if cap is not None: - s_i_j = jnp.tanh(s_i_j / cap) - s_i_j = s_i_j * cap - s_i_j = jnp.where(broadcasted_mask, s_i_j, mask_value) - m_i_j = s_i_j.max(axis=-1) - p_i_j = jnp.exp(s_i_j - m_i_j[..., None]) + s_i_j_dup = jnp.tanh(s_i_j_dup / cap) + s_i_j_dup = s_i_j_dup * cap + m_i_j = s_i_j_dup.max(axis=-1) + + def fuse_this(q_slice, k_j_slice, mask_value, broadcasted_mask, m_i_j): + s_i_j = jnp.einsum( + "bxhqc,bxkc->bxhqk", + q_slice, + k_j_slice, + preferred_element_type=jnp.bfloat16, + ) + if unroll: + if i == j: + s_i_j = jnp.where(broadcasted_mask, s_i_j, mask_value) + else: + s_i_j = jnp.where(broadcasted_mask, s_i_j, mask_value) + if cap is not None: + s_i_j = jnp.tanh(s_i_j / cap) + s_i_j = s_i_j * cap + p_i_j = jnp.exp(s_i_j - m_i_j[..., None]) + return p_i_j + + p_i_j = must_fuse_call("1")(fuse_this)( + q_slice, k_j_slice, mask_value, broadcasted_mask, m_i_j + ) l_i_j = p_i_j.sum(axis=-1) assert m_i_j.shape == m_i_slice.shape m_i_new = jnp.maximum(m_i_slice, m_i_j) @@ -173,19 +226,35 @@ def compute_attention_block(output, l, m): l_i_new = m_i_difference * l_i_slice + m_i_j_difference * l_i_j divider = l_i_new[..., None] - numerator = l_i_slice[..., None] * m_i_difference[..., None] * output_i_slice + m_i_j_difference[ - ..., None - ] * jnp.einsum( + pv = jnp.einsum( "bxhqk,bxkc->bxhqc", p_i_j, v_j_slice, preferred_element_type=data_type, ) + # This forces the layout of the final @V to have better utilization by + # avoiding using XLU: + # Changes conv emitter type from + # EmitAllInputFeatureInSublanesOutputBatchInSublanesXposeReuse to + # EmitInputBatchInLanes + pv = layout.with_layout_constraint( + pv, DLL(major_to_minor=(0, 1, 2, 4, 3)) + ) + numerator = ( + l_i_slice[..., None] * m_i_difference[..., None] * output_i_slice + + m_i_j_difference[..., None] * pv + ) output_i_slice_new = numerator / divider - output = jax.lax.dynamic_update_index_in_dim(output, output_i_slice_new, i * block_q, axis=-2) - l = jax.lax.dynamic_update_index_in_dim(l, l_i_new, i * block_q, axis=-1) - m = jax.lax.dynamic_update_index_in_dim(m, m_i_new, i * block_q, axis=-1) + output = jax.lax.dynamic_update_index_in_dim( + output, output_i_slice_new, i * block_q, axis=-2 + ) + l = jax.lax.dynamic_update_index_in_dim( + l, l_i_new, i * block_q, axis=-1 + ) + m = jax.lax.dynamic_update_index_in_dim( + m, m_i_new, i * block_q, axis=-1 + ) return output, l, m def identity(output, l, m): @@ -193,27 +262,40 @@ def identity(output, l, m): return output, l, m - batch_size = mask_blocked.shape[0] - mask_i_j_slice = jax.lax.dynamic_slice(mask_blocked, (0, i, j), (batch_size, 1, 1)) - # The compute_attention_block should be executed if at least one element - # in the slice is non-zero, meaning at least one batch requires work for - # this block. - output, l, m = jax.lax.cond( - jnp.any(jnp.not_equal(mask_i_j_slice, 0)), - compute_attention_block, - identity, - output, - l, - m, - ) + if unroll: + if should_compute_block(i, j): + output, l, m = compute_attention_block(output, l, m) + else: + output, l, m = identity(output, l, m) + else: + output, l, m = jax.lax.cond( + should_compute_block(i, j), + compute_attention_block, + identity, + output, + l, + m, + ) return output, l, m - output, l, m = jax.lax.fori_loop(0, num_q_blocks, inner_loop_body, (output, l, m), unroll=True) + if unroll: + for i in range(num_q_blocks): + output, l, m = inner_loop_body(i, (output, l, m)) + else: + output, l, m = jax.lax.fori_loop( + 0, num_q_blocks, inner_loop_body, (output, l, m), unroll=False + ) return (output, l, m) - output, l, m = jax.lax.fori_loop(0, num_kv_blocks, outer_loop_body, (output, l, m), unroll=True) + if unroll: + for j in range(num_kv_blocks): + output, l, m = outer_loop_body(j, (output, l, m)) + else: + output, l, m = jax.lax.fori_loop( + 0, num_kv_blocks, outer_loop_body, (output, l, m), unroll=False + ) # Reshape the output to drop the size one dimension at index 2, # which corresponds to `num_q_heads // num_kv_heads` when