Skip to content
Draft
Show file tree
Hide file tree
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
8 changes: 4 additions & 4 deletions .github/workflows/ci_pipeline.yml
Original file line number Diff line number Diff line change
Expand Up @@ -86,12 +86,12 @@ jobs:

build_and_upload_maxtext_package:
name: Build MaxText Package
needs: [analyze_code_changes, code_quality_check, docs_build_check]
# Run if either tests or notebooks need to run; on PRs, gate on code quality + docs passing
needs: [analyze_code_changes, code_quality_check]
# Run if either tests or notebooks need to run; on PRs, gate on code quality
if: |
always() &&
(needs.analyze_code_changes.outputs.run_tests == 'true' || needs.analyze_code_changes.outputs.run_notebooks == 'true') &&
(github.event_name != 'pull_request' || (needs.code_quality_check.result == 'success' && needs.docs_build_check.result == 'success'))
(github.event_name != 'pull_request' || needs.code_quality_check.result == 'success')
uses: ./.github/workflows/build_package.yml
with:
device_type: tpu
Expand Down Expand Up @@ -184,7 +184,7 @@ jobs:
if: |
always() &&
needs.gate_test_run.result == 'success' &&
github.ref == 'refs/heads/main' && (github.event_name == 'schedule' || github.event_name == 'workflow_dispatch')
(github.event_name == 'schedule' || github.event_name == 'workflow_dispatch')
uses: ./.github/workflows/run_tests_coordinator.yml
strategy:
fail-fast: false
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -2094,9 +2094,6 @@ def _make_dynamic_splash_attention(
if config is None:
config = SplashConfig.get_default()

# This is the only mode that supports the dynamic grid.
config = dataclasses.replace(config, dq_reduction_steps=3)

def process_mask_shard(mask):
process_mask_fn = functools.partial(
mask_info_lib._process_dynamic_mask,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -300,74 +300,34 @@ def _process_dynamic_mask(
raise ValueError(f"{kv_block_size=} should divide {kv_seq_len=}.")

# Tile the last 2 dimensions of the mask into 2D tiles of size `block_shape`.
mask_blocks = (
mask.reshape(
q_blocks_count,
q_block_size,
kv_blocks_count,
kv_block_size,
)
.swapaxes(-2, -3)
.astype(partial_mask_blocks_dtype)
)

any_mask = jnp.any(mask_blocks, axis=(-1, -2)).astype(np.int32)
all_mask = jnp.all(mask_blocks, axis=(-1, -2)).astype(np.int32)
block_mask = any_mask + all_mask
mask_blocks = mask.reshape(
q_blocks_count,
q_block_size,
kv_blocks_count,
kv_block_size,
).swapaxes(1, 2)

block_ids = jnp.arange(block_mask.size, dtype=np.int32).reshape(block_mask.shape)
if is_dkv:
block_mask = block_mask.swapaxes(-1, -2)
block_ids = block_ids.swapaxes(-1, -2)
mask_blocks = mask_blocks.swapaxes(0, 1)
mask_blocks = mask_blocks.swapaxes(-1, -2)

active_mask = block_mask > 0
# If an entire row is masked then that output tile won't be visited.
# We extend the grid to visit these tiles to initialize them.
empty_rows = jnp.all(block_mask == 0, axis=-1)
first_col = jnp.arange(block_mask.shape[1]) == 0
active_mask |= empty_rows[:, None] & first_col

num_active_blocks = active_mask.flatten().sum(keepdims=True)
active_indices = jnp.argwhere(active_mask, size=active_mask.size, fill_value=-1)
active_rows = active_indices[:, 0].astype(np.int32)
active_cols = active_indices[:, 1].astype(np.int32)

block_mask = block_mask[active_rows, active_cols]
mask_next = block_ids.at[active_rows, active_cols].get(wrap_negative_indices=False)
mask_next = jnp.where(block_mask == 1, mask_next, 0)

# Mask out the blocks that aren't active.
mask = (jnp.arange(block_mask.size) < num_active_blocks).astype(np.int32)
block_mask = block_mask * mask

# Collapsing because the block ids are linearized.
mask_blocks = lax.collapse(mask_blocks, 0, 2)

def _downcast(array: jax.Array, max_value: int) -> jax.Array:
if array.size == 0:
return array

if array.dtype != np.int32:
raise ValueError(f"Expected int32 input, but got {array.dtype}.")

if max_value <= np.iinfo(np.int8).max:
return array.astype(np.int8)
elif max_value <= np.iinfo(np.int16).max:
return array.astype(np.int16)
else:
return array.astype(np.int32)
num_blocks = q_blocks_count * kv_blocks_count
mask_blocks = mask_blocks.reshape(num_blocks, mask_blocks.shape[-2], mask_blocks.shape[-1])
mask_blocks = mask_blocks.astype(partial_mask_blocks_dtype)

mask_next = jnp.arange(num_blocks, dtype=jnp.int32)
if downcast_smem_data:
block_mask = block_mask.astype(np.int8) # values are in the range [0, 1, 2]
mask_next = _downcast(mask_next, q_blocks_count * kv_blocks_count)
if num_blocks <= np.iinfo(np.int8).max:
mask_next = mask_next.astype(np.int8)
elif num_blocks <= np.iinfo(np.int16).max:
mask_next = mask_next.astype(np.int16)

return MaskInfo(
mask_next=mask_next,
active_rows=active_rows,
active_cols=active_cols,
block_mask=block_mask,
num_active_blocks=num_active_blocks,
active_rows=None,
active_cols=None,
block_mask=None,
num_active_blocks=None,
partial_mask_blocks=mask_blocks,
q_sequence=None,
kv_sequence=None,
Expand Down
54 changes: 46 additions & 8 deletions src/maxtext/layers/attention_op.py
Original file line number Diff line number Diff line change
Expand Up @@ -2108,11 +2108,38 @@ def wrap_flash_attention(
decoder_segment_ids_tuple = None

if self.config.use_tokamax_splash:
orig_q_len = query.shape[2]
orig_kv_len = key.shape[2]
padded_q_len = ((orig_q_len + sa_config.block_q - 1) // sa_config.block_q) * sa_config.block_q
padded_kv_len = ((orig_kv_len + sa_config.block_kv - 1) // sa_config.block_kv) * sa_config.block_kv
pad_q = padded_q_len - orig_q_len
pad_kv = padded_kv_len - orig_kv_len

if pad_q > 0:
query = jnp.pad(query, ((0, 0), (0, 0), (0, pad_q), (0, 0)))
if decoder_segment_ids_tuple is not None:
decoder_segment_ids_q = jnp.pad(decoder_segment_ids_tuple.q, ((0, 0), (0, pad_q)), constant_values=-1)
else:
if decoder_segment_ids_tuple is not None:
decoder_segment_ids_q = decoder_segment_ids_tuple.q

if pad_kv > 0:
key = jnp.pad(key, ((0, 0), (0, 0), (0, pad_kv), (0, 0)))
value = jnp.pad(value, ((0, 0), (0, 0), (0, pad_kv), (0, 0)))
if decoder_segment_ids_tuple is not None:
decoder_segment_ids_kv = jnp.pad(decoder_segment_ids_tuple.kv, ((0, 0), (0, pad_kv)), constant_values=-1)
else:
if decoder_segment_ids_tuple is not None:
decoder_segment_ids_kv = decoder_segment_ids_tuple.kv

if decoder_segment_ids_tuple is not None:
decoder_segment_ids_tuple = tokamax_splash_kernel.SegmentIds(decoder_segment_ids_q, decoder_segment_ids_kv)

segment_in_axis = 0 if decoder_segment_ids_tuple is not None else None

if indexer_mask is not None:
# Convert additive float mask (0.0=attend, negative=masked) to boolean mask for Tokamax splash kernel
indexer_mask = indexer_mask == 0.0
pad_q = mask_shape[0] - indexer_mask.shape[-2]
pad_kv = mask_shape[1] - indexer_mask.shape[-1]
indexer_mask = indexer_mask >= DEFAULT_MASK_VALUE * 0.5
if pad_q > 0 or pad_kv > 0:
pad_width = [(0, 0)] * (indexer_mask.ndim - 2) + [(0, pad_q), (0, pad_kv)]
indexer_mask = jnp.pad(
Expand All @@ -2136,13 +2163,18 @@ def dynamic_mask_splash_kernel(q, k, v, segment, sinks, indexer_mask):
return kernel(q, k, v, segment, sinks=sinks), None

# Iterate over batch dimension for (query, key, value, segment, sinks, mask)
attn_fn = jax.vmap(dynamic_mask_splash_kernel, (0, 0, 0, 0, None, 0))
attn_fn = jax.vmap(dynamic_mask_splash_kernel, (0, 0, 0, segment_in_axis, None, 0))

if record_max_logits:
attention_output, max_logits = attn_fn(query, key, value, decoder_segment_ids_tuple, sinks, indexer_mask)
if pad_q > 0:
attention_output = attention_output[:, :, :orig_q_len, :]
max_logits = max_logits[:, :, :orig_q_len]
return attention_output, max_logits
else:
attention_output, _ = attn_fn(query, key, value, decoder_segment_ids_tuple, sinks, indexer_mask)
if pad_q > 0:
attention_output = attention_output[:, :, :orig_q_len, :]
return attention_output, None
else:
kernel = partial(splash_kernel, max_logit_value=max_logit_value)
Expand All @@ -2154,14 +2186,20 @@ def kernel_fn(q, k, v, d, s):
out, stats = kernel(q, k, v, d, sinks=s, save_residuals=True)
return out, stats["max_logits"]

attention_output, max_logits = jax.vmap(kernel_fn, in_axes=(0, 0, 0, 0, None))(
attention_output, max_logits = jax.vmap(kernel_fn, in_axes=(0, 0, 0, segment_in_axis, None))(
query, key, value, decoder_segment_ids_tuple, sinks
)
if pad_q > 0:
attention_output = attention_output[:, :, :orig_q_len, :]
max_logits = max_logits[:, :, :orig_q_len]
return attention_output, max_logits
else:
attention_output = jax.vmap(lambda q, k, v, d, s: kernel(q, k, v, d, sinks=s), in_axes=(0, 0, 0, 0, None))(
query, key, value, decoder_segment_ids_tuple, sinks
)
attention_output = jax.vmap(
lambda q, k, v, d, s: kernel(q, k, v, d, sinks=s),
in_axes=(0, 0, 0, segment_in_axis, None),
)(query, key, value, decoder_segment_ids_tuple, sinks)
if pad_q > 0:
attention_output = attention_output[:, :, :orig_q_len, :]
return attention_output, None

elif self.config.use_jax_splash:
Expand Down
Loading