Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
162 changes: 122 additions & 40 deletions src/maxtext/kernels/attention/jax_flash_attention.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand All @@ -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.

Expand All @@ -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:
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand All @@ -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)
Expand All @@ -173,47 +226,76 @@ 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):
"""A no-op identity function."""

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
Expand Down
Loading