Skip to content
Open
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
9 changes: 9 additions & 0 deletions src/maxtext/common/metric_logger.py
Original file line number Diff line number Diff line change
Expand Up @@ -223,6 +223,10 @@ def _log_training_metrics(self, metrics, step):
log_parts.append(f"main_model_loss: {loss - mtp_loss:.3f}")
log_parts.append(f"mtp_loss: {mtp_loss:.3f}")

if getattr(self.config, "use_indexer", False):
indexer_l = scalars.get("learning/indexer_loss", 0.0)
log_parts.append(f"indexer_loss: {indexer_l:.3f}")

max_logging.log(", ".join(log_parts))

def _log_eval_metrics(self, metrics, step):
Expand All @@ -246,6 +250,11 @@ def _log_eval_metrics(self, metrics, step):
)
if "eval/avg_dpo_reward_accuracy" in scalars:
log_parts.append(f"dpo_reward_accuracy={scalars['eval/avg_dpo_reward_accuracy']:.3f}")

if getattr(self.config, "use_indexer", False):
indexer_l = scalars.get("eval/avg_indexer_loss", 0.0)
log_parts.append(f"avg_indexer_loss={indexer_l:.3f}")

max_logging.log(", ".join(log_parts))

def _log_running_eval_metrics(self, metrics, step):
Expand Down
2 changes: 0 additions & 2 deletions src/maxtext/configs/types.py
Original file line number Diff line number Diff line change
Expand Up @@ -3650,8 +3650,6 @@ def calculate_global_batch_sizes(per_device_batch_size, expansion_factor, num_de
raise ValueError("TPU Tokamax ring attention does not support ragged attention.")
if self.attention_sink:
raise ValueError("TPU Tokamax ring attention does not support attention sinks.")
if self.use_indexer:
raise ValueError("TPU Tokamax ring attention does not support sparse indexer masks.")
if self.use_chunked_prefill:
raise ValueError("TPU Tokamax ring attention does not support chunked prefill yet.")
if self.moba:
Expand Down
41 changes: 29 additions & 12 deletions src/maxtext/kernels/attention/tokamax_ring_attention.py
Original file line number Diff line number Diff line change
Expand Up @@ -132,8 +132,6 @@ def validate_tokamax_ring_runtime(
raise ValueError("TPU Tokamax ring attention does not support chunked prefill yet.")
if sinks is not None:
raise ValueError("TPU Tokamax ring attention does not support attention sinks.")
if indexer_mask is not None:
raise ValueError("TPU Tokamax ring attention does not support indexer masks.")
if bidirectional_mask is not None:
raise ValueError("TPU Tokamax ring attention does not support bidirectional masks.")
if record_max_logits:
Expand Down Expand Up @@ -304,6 +302,7 @@ def make_sharded_ring_attention_kernel(
ring_axis: str,
attn_logits_soft_cap: float | None,
maybe_shard_with_pspec: Any,
mask: Any = None,
):
"""Builds and shards the Tokamax ring attention kernel for MaxText."""
splash_config = build_splash_config(
Expand All @@ -316,11 +315,17 @@ def make_sharded_ring_attention_kernel(
if config.use_max_logit_estimate > 0:
splash_config = dataclasses.replace(splash_config, max_logit_const=config.use_max_logit_estimate)

mask = _make_causal_mask(
(query.shape[2], key.shape[2]),
context_parallel_size,
load_balanced=config.context_parallel_load_balance,
)
if mask is None:
# When using the indexer, causal masking is unified into the dynamic indexer_mask
# and applied dynamically per block; use FullMask to avoid duplicate static masks.
if getattr(config, "use_indexer", False):

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Can we raise when use_indexer is set but indexer_mask is None? If the indexer returns None then would be no causal at all

mask = tokamax_splash_mask.FullMask((query.shape[2], key.shape[2]))
else:
mask = _make_causal_mask(
(query.shape[2], key.shape[2]),
context_parallel_size,
load_balanced=config.context_parallel_load_balance,
)

@functools.partial(jax.jit, static_argnames=["single_head_mask"])
def wrap_ring_kernel(single_head_mask):
Expand Down Expand Up @@ -352,15 +357,27 @@ def call_ring_attention(
decoder_segment_ids_q: Any,
decoder_segment_ids_kv: Any,
ring_kernel: Any,
indexer_mask: Any = None,
):
"""Calls a Tokamax ring attention kernel over the MaxText batch dimension."""
if (decoder_segment_ids_q is None) != (decoder_segment_ids_kv is None):
raise ValueError("decoder_segment_ids_q and decoder_segment_ids_kv must both be set or both be None.")
# Vectorize execution across batch dimension, threading indexer_mask when present.
# Note: ring_kernel expects positional arguments (q, k, v, segment_ids, sinks, indexer_mask).
if decoder_segment_ids_q is None:
return jax.vmap(lambda q, k, v: ring_kernel(q, k, v, None), in_axes=(0, 0, 0))(query, key, value)

def call_one(q, k, v, q_segment_ids, kv_segment_ids):
if indexer_mask is None:
return jax.vmap(lambda q, k, v: ring_kernel(q, k, v, None, None, None), in_axes=(0, 0, 0))(query, key, value)
return jax.vmap(
lambda q, k, v, im: ring_kernel(q, k, v, None, None, im),
in_axes=(0, 0, 0, 0),
)(query, key, value, indexer_mask)

def call_one(q, k, v, q_segment_ids, kv_segment_ids, im=None):
segment_ids = ring_attention_kernel.SegmentIds(q_segment_ids, kv_segment_ids)
return ring_kernel(q, k, v, segment_ids)
return ring_kernel(q, k, v, segment_ids, None, im)

return jax.vmap(call_one, in_axes=(0, 0, 0, 0, 0))(query, key, value, decoder_segment_ids_q, decoder_segment_ids_kv)
if indexer_mask is None:
return jax.vmap(call_one, in_axes=(0, 0, 0, 0, 0))(query, key, value, decoder_segment_ids_q, decoder_segment_ids_kv)
return jax.vmap(call_one, in_axes=(0, 0, 0, 0, 0, 0))(
query, key, value, decoder_segment_ids_q, decoder_segment_ids_kv, indexer_mask
)
145 changes: 124 additions & 21 deletions src/maxtext/kernels/tokamax_splash_attention/ring_attention_kernel.py
Original file line number Diff line number Diff line change
Expand Up @@ -60,8 +60,61 @@ def _validate_ring_axis_size(ring_axis: str, ring_axis_size: int, expected_ring_
)


def _inject_local_indexer_mask(
local_mask_info: MaskInfo,
local_idx_mask: jax.Array | None,
block_shape: tuple[int, int] = (128, 128),
is_dkv: bool = False,
) -> MaskInfo:
"""Injects a pre-sliced dynamic Indexer mask shard into local MaskInfo for the current ring step."""
if local_idx_mask is None:
return local_mask_info

bq, bkv = block_shape
if local_idx_mask.ndim == 3:
local_idx_mask = local_idx_mask[0]
if local_idx_mask.dtype != jnp.bool_:
local_idx_mask = jnp.isclose(local_idx_mask, 0.0)

q_len, kv_len = local_idx_mask.shape
# Since causal and padding masks are already fully integrated into the global indexer_mask
# tensor (via indexer_mask += attention_mask in attention_mla.py) before slicing,
# local_idx_mask already contains the complete causally-masked top-k selection for this ring hop.
combined_mask = local_idx_mask

# Tile 2D mask into hardware block chunks [bq, bkv]
q_blocks = q_len // bq
kv_blocks = kv_len // bkv
Comment thread
zcjhao marked this conversation as resolved.
num_blocks = q_blocks * kv_blocks

blocks = combined_mask.reshape(q_blocks, bq, kv_blocks, bkv)
blocks = blocks.swapaxes(1, 2) # [q_blocks, kv_blocks, bq, bkv]

if is_dkv:
# SplashAttention dkv grids are scheduled as KV-major (kv_blocks, q_blocks).
# We transpose both block grid and intra-block axes to match Pallas grid_idx order.
blocks = blocks.swapaxes(0, 1) # [kv_blocks, q_blocks, bq, bkv]
blocks = blocks.swapaxes(-1, -2) # [kv_blocks, q_blocks, bkv, bq]

blocks = blocks.reshape(num_blocks, blocks.shape[-2], blocks.shape[-1])
blocks = blocks.astype(jnp.int8)

mask_next = jnp.arange(num_blocks, dtype=jnp.int32)
return local_mask_info._replace(
mask_next=mask_next,
active_rows=None,
active_cols=None,
block_mask=None,
num_active_blocks=None,
partial_mask_blocks=blocks,
q_sequence=None,
kv_sequence=None,
)


def _ring_attention_forward(
fwd_mask_info: MaskInfo,
indexer_mask: jax.Array | None,
q: jax.Array,
k: jax.Array,
v: jax.Array,
Expand Down Expand Up @@ -118,12 +171,33 @@ def _ring_attention_forward(
l_init = jnp.zeros((o_shape[0], o_shape[1]), jnp.float32)
m_init = jnp.full_like(l_init, mask_value, dtype=jnp.float32)

def body(carry, i: int):
m_prev, l_prev, o_prev, k_current, v_current, segment_ids_current = carry
if indexer_mask is not None:
# Reshape global indexer mask to [..., ring_axis_size, kv_shard_len] for dynamic step slicing.
kv_shard_len = k.shape[-2]
mask_4d = indexer_mask.reshape(*indexer_mask.shape[:-1], ring_axis_size, kv_shard_len)
else:
mask_4d = None

xs = jnp.arange(0, ring_axis_size)

def body(carry, i):
m_prev, l_prev, o_prev, k_current, v_current, segment_ids_current = carry
current_kv_shard_idx = (ring_axis_idx - i) % ring_axis_size
if mask_4d is not None:
# Slice EXACTLY the current KV shard's mask block.
local_idx_mask = jax.lax.dynamic_slice_in_dim(mask_4d, current_kv_shard_idx, 1, axis=-2)
local_idx_mask = jnp.squeeze(local_idx_mask, axis=-2)
else:
local_idx_mask = None

local_fwd_mask_info = _dynamic_slice_mask_info(fwd_mask_info, current_kv_shard_idx, ring_axis_size)
local_fwd_mask_info = _offset_q_sequence_for_kv_shard(local_fwd_mask_info, current_kv_shard_idx, k_current.shape[-2])
local_fwd_mask_info = _inject_local_indexer_mask(
local_fwd_mask_info,
local_idx_mask,
block_shape=(config.block_q, config.block_kv),
is_dkv=False,
)
k_next = shift(k_current)
v_next = shift(v_current)

Expand Down Expand Up @@ -168,7 +242,7 @@ def body(carry, i: int):
(m_final, l_final, o_final, _, _, _), _ = lax.scan(
body,
initial_carry,
xs=jnp.arange(0, ring_axis_size),
xs=xs,
length=ring_axis_size,
unroll=config.ring_scan_unroll,
) # type: ignore[arg-type]
Expand Down Expand Up @@ -198,7 +272,7 @@ def _ring_attention_bwd(
do: jax.Array,
):
del save_residuals
(q, k, v, segment_ids, sinks, out, logsumexp, dkv_mask_info) = res
(q, k, v, segment_ids, sinks, out, logsumexp, dkv_mask_info, indexer_mask) = res
do = do.astype(jnp.float32)
if dkv_mask_info is None:
raise ValueError("Need to specify backward blocks.")
Expand Down Expand Up @@ -229,10 +303,29 @@ def rotate_kv(k_current, v_current, segment_ids_current):
segment_ids_next = None
return k_next, v_next, segment_ids_next

def compute_step(i: int, k_current, v_current, segment_ids_current, dq_accum):
if indexer_mask is not None:
# Reshape global indexer mask to [..., ring_axis_size, kv_shard_len] for backward rotation slicing.
kv_shard_len = k.shape[-2]
mask_4d = indexer_mask.reshape(*indexer_mask.shape[:-1], ring_axis_size, kv_shard_len)
step0_mask_idx = (ring_axis_idx - 0) % ring_axis_size
step0_mask = jax.lax.dynamic_slice_in_dim(mask_4d, step0_mask_idx, 1, axis=-2)
step0_mask = jnp.squeeze(step0_mask, axis=-2)
else:
mask_4d = None
step0_mask = None

xs = jnp.arange(1, ring_axis_size)

def compute_step(i: int, local_idx_mask, k_current, v_current, segment_ids_current, dq_accum):
current_kv_shard_idx = (ring_axis_idx - i) % ring_axis_size
local_dkv_mask_info = _dynamic_slice_mask_info(dkv_mask_info, current_kv_shard_idx, ring_axis_size)
local_dkv_mask_info = _offset_q_sequence_for_kv_shard(local_dkv_mask_info, current_kv_shard_idx, k_current.shape[-2])
local_dkv_mask_info = _inject_local_indexer_mask(
local_dkv_mask_info,
local_idx_mask,
block_shape=(config.block_q_dkv, config.block_kv_dkv),
is_dkv=True,
)

residuals_for_chunk = (
q,
Expand Down Expand Up @@ -266,11 +359,17 @@ def compute_step(i: int, k_current, v_current, segment_ids_current, dq_accum):
dq_i = dq_accum + dq_i.astype(jnp.float32)
return dq_i, dk_i, dv_i, dsinks

dq_i, dk_pending, dv_pending, dsinks = compute_step(0, k, v, segment_ids, dq_accum)
dq_i, dk_pending, dv_pending, dsinks = compute_step(0, step0_mask, k, v, segment_ids, dq_accum)
dq_accum = dq_i
k_current, v_current, segment_ids_current = rotate_kv(k, v, segment_ids)

def body(carry, i: int):
def body(carry, i):
if mask_4d is not None:
current_kv_shard_idx = (ring_axis_idx - i) % ring_axis_size
local_idx_mask = jax.lax.dynamic_slice_in_dim(mask_4d, current_kv_shard_idx, 1, axis=-2)
local_idx_mask = jnp.squeeze(local_idx_mask, axis=-2)
else:
local_idx_mask = None
(
dq_accum,
dk_accum,
Expand All @@ -285,7 +384,7 @@ def body(carry, i: int):
dk_next = shift(dk_accum + dk_pending.astype(jnp.float32))
dv_next = shift(dv_accum + dv_pending.astype(jnp.float32))
k_next, v_next, segment_ids_next = rotate_kv(k_current, v_current, segment_ids_current)
dq_i, dk_i, dv_i, dsinks = compute_step(i, k_current, v_current, segment_ids_current, dq_accum)
dq_i, dk_i, dv_i, dsinks = compute_step(i, local_idx_mask, k_current, v_current, segment_ids_current, dq_accum)
dq_accum = dq_i
return (
dq_accum,
Expand Down Expand Up @@ -313,7 +412,7 @@ def body(carry, i: int):
(dq, dk, dv, dk_pending, dv_pending, _, _, _, dsinks), _ = lax.scan(
body,
initial_carry,
xs=jnp.arange(1, ring_axis_size),
xs=xs,
length=ring_axis_size - 1,
unroll=config.ring_scan_unroll,
)
Expand All @@ -333,6 +432,7 @@ def body(carry, i: int):
dv.astype(v.dtype),
None,
dsinks,
None, # indexer_mask
)


Expand All @@ -344,6 +444,7 @@ def _ring_attention_fwd(
v: jax.Array,
segment_ids: SegmentIds | None,
sinks: jax.Array | None,
indexer_mask: jax.Array | None,
# nondiff_args
mask_value: float, # 1
is_mqa: bool, # 2
Expand Down Expand Up @@ -388,6 +489,7 @@ def _ring_attention_fwd(

out, (logsumexp, max_logits) = _ring_attention_forward(
fwd_mask_info,
indexer_mask,
q,
k,
v,
Expand All @@ -405,7 +507,7 @@ def _ring_attention_fwd(
if config.residual_checkpoint_name is not None:
out = ad_checkpoint.checkpoint_name(out, name=config.residual_checkpoint_name)
logsumexp = ad_checkpoint.checkpoint_name(logsumexp, name=config.residual_checkpoint_name)
residuals = (q, k, v, segment_ids, sinks, out, logsumexp, dkv_mask_info)
residuals = (q, k, v, segment_ids, sinks, out, logsumexp, dkv_mask_info, indexer_mask)
return out, residuals


Expand All @@ -431,6 +533,7 @@ def _ring_attention_custom(
v: jax.Array,
segment_ids: SegmentIds | None,
sinks: jax.Array | None,
indexer_mask: jax.Array | None,
mask_value: float,
is_mqa: bool,
config: SplashConfig,
Expand Down Expand Up @@ -468,6 +571,7 @@ def _ring_attention_custom(
del dkv_mask_info, dkv_mask_sparsity
out, _ = _ring_attention_forward(
fwd_mask_info,
indexer_mask,
q,
k,
v,
Expand Down Expand Up @@ -509,6 +613,7 @@ def _ring_attention(
v: jax.Array,
segment_ids: SegmentIds | None = None,
sinks: jax.Array | None = None,
indexer_mask: jax.Array | None = None,
*,
is_mqa: bool,
config: SplashConfig,
Expand Down Expand Up @@ -559,6 +664,7 @@ def _ring_attention(
v,
segment_ids,
sinks,
indexer_mask,
is_mqa=is_mqa,
config=config,
mask_value=mask_value,
Expand Down Expand Up @@ -637,19 +743,16 @@ def mask_info_spec(mask_info):
if mask_info is None:
return None
return MaskInfo( # pytype: disable=wrong-arg-types
mask_next=_resolve_spec(mask_info.mask_next), # pyrefly: ignore[bad-argument-type]
active_rows=_resolve_spec(mask_info.active_rows), # pyrefly: ignore[bad-argument-type]
active_cols=_resolve_spec(mask_info.active_cols), # pyrefly: ignore[bad-argument-type]
num_active_blocks=_resolve_spec(mask_info.num_active_blocks), # pyrefly: ignore[bad-argument-type]
block_mask=_resolve_spec(mask_info.block_mask), # pyrefly: ignore[bad-argument-type]
partial_mask_blocks=jax.sharding.PartitionSpec() # replicated # pyrefly: ignore[bad-argument-type]
mask_next=_resolve_spec(mask_info.mask_next),
active_rows=_resolve_spec(mask_info.active_rows),
active_cols=_resolve_spec(mask_info.active_cols),
num_active_blocks=_resolve_spec(mask_info.num_active_blocks),
block_mask=_resolve_spec(mask_info.block_mask),
partial_mask_blocks=jax.sharding.PartitionSpec() # replicated
if mask_info.partial_mask_blocks is not None
else None,
q_sequence=_resolve_spec(mask_info.q_sequence), # pyrefly: ignore[bad-argument-type]
# pyrefly: ignore[bad-argument-type]
kv_sequence=jax.sharding.PartitionSpec()
if mask_info.kv_sequence is not None
else None, # pyrefly: ignore[bad-argument-type]
q_sequence=_resolve_spec(mask_info.q_sequence),
kv_sequence=jax.sharding.PartitionSpec() if mask_info.kv_sequence is not None else None,
)

return RingSplashAttentionKernel(
Expand Down
Loading
Loading