Skip to content

Commit af8a014

Browse files
committed
[Memory] Implement QK attention head chunking to prevent HBM OOM
Currently, evaluating the QK dot product natively in MLA materializes the full [batch, count_of_heads, q_len, kv_len] attention scores tensor. This causes severe memory constraints (HBM OOM) for large context lengths. This change introduces the Config variable `qk_head_chunk_size` alongside a `jax.lax.scan` algorithm. Since the `heads` dimension is locally dense and unsharded (unlike `q_len` which is Context Parallel), we scan over groups of heads iteratively computing query-key projections and their softmax partials, avoiding vast intermediate allocations.
1 parent d557e88 commit af8a014

3 files changed

Lines changed: 114 additions & 22 deletions

File tree

src/maxtext/configs/base.yml

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1336,3 +1336,6 @@ elastic_enabled: false
13361336
elastic_timeout_seconds: 300
13371337
elastic_max_retries: 10
13381338
elastic_min_slice_count: -1
1339+
1340+
# Limits HBM footprint by sequentially evaluating the QK matrix across the unsharded local heads dimension natively. Size must evenly divide the number of attention heads.
1341+
qk_head_chunk_size: 0

src/maxtext/configs/types.py

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -637,6 +637,14 @@ class Attention(BaseModel):
637637
force_q_layout: bool = Field(False, description="Force the Q layout")
638638
use_qk_clip: bool = Field(False, description="Whether to use QK-Clip (MuonClip) for training stability.")
639639
qk_clip_threshold: float = Field(100.0, description="Threshold for QK-Clip (tau).")
640+
qk_head_chunk_size: int = Field(
641+
0,
642+
description=(
643+
"Chunk size over heads dimension for QK attention dot product. "
644+
"Default is 0 (no chunking). Must divide the number of heads evenly. "
645+
"This reduces memory footprint at the cost of time, especially helpful in long context."
646+
),
647+
)
640648

641649

642650
class MoBa(BaseModel):

src/maxtext/layers/attention_mla.py

Lines changed: 103 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -361,18 +361,56 @@ def __call__(
361361
if k.shape[1] <= self.indexer_topk:
362362
return None, None, None
363363

364-
# Compute Index Scores
365-
# QK product: relu(q @ k.T), [b, t, s, h]
366-
# Similar to MQA, each key is shared by h query head
367-
logits = jnp.einsum("bthd, bsd -> btsh", q, k, precision=self.config.matmul_precision)
368-
logits = jax.nn.relu(logits)
369364
# Compute head weights: project from input, [b, t, embed_dim] -> [b, t, h]
370365
weights = self.weights_proj(inputs_q)
371366
# Weights scaling affect indexer_score, but does not affect topk_indices. Keep scaling for numerical stability.
372367
# https://github.com/deepseek-ai/DeepSeek-V3.2-Exp/blob/87e509a2e5a100d221c97df52c6e8be7835f0057/inference/model.py#L478-L480
373368
weights = weights * (self.n_heads**-0.5) * self.softmax_scale
374-
# Aggregate head-wise logits: logits @ weights
375-
indexer_score = jnp.einsum("btsh, bth -> bts", logits, weights, precision=self.config.matmul_precision) # [b, t, s]
369+
370+
# Compute Index Scores
371+
# When qk_head_chunk_size > 0, compute Index Scores by chunking the 'heads' dimension to reduce memory
372+
# The naive evaluation materializes [b, t, s, h].
373+
# We use jax.lax.scan to compute the score iteratively over head chunks.
374+
b, t, h, d = q.shape
375+
# Control the HBM footprint of QK tensor: [batch, q_len, s_len, heads]
376+
# If set to 0 (defaults), it falls back to native materialization.
377+
head_chunk_size = getattr(self.config, "qk_head_chunk_size", 0)
378+
if head_chunk_size > 0:
379+
if head_chunk_size > h or h % head_chunk_size != 0:
380+
raise ValueError(
381+
f"qk_head_chunk_size ({head_chunk_size}) must be <= number of heads ({h}) " f"and divide it evenly."
382+
)
383+
num_chunks = h // head_chunk_size
384+
# q: [b, t, h, d] -> [h, b, t, d] -> [num_chunks, head_chunk_size, b, t, d]
385+
q_h = q.transpose(2, 0, 1, 3).reshape(num_chunks, head_chunk_size, b, t, d)
386+
# weights: [b, t, h] -> [h, b, t] -> [num_chunks, head_chunk_size, b, t]
387+
w_h = weights.transpose(2, 0, 1).reshape(num_chunks, head_chunk_size, b, t)
388+
389+
def scan_body_indexer(carry, xs):
390+
q_c = xs["q"] # [h_chunk, b, t, d]
391+
w_c = xs["w"] # [h_chunk, b, t]
392+
393+
# Directly use the chunked shapes in einsum to avoid transposes inside the loop
394+
logits = jnp.einsum("hbtd, bsd -> btsh", q_c, k, precision=self.config.matmul_precision)
395+
logits = jax.nn.relu(logits)
396+
397+
score_chunk = jnp.einsum(
398+
"btsh, hbt -> bts",
399+
logits,
400+
w_c,
401+
precision=self.config.matmul_precision,
402+
)
403+
return carry + score_chunk.astype(jnp.float32), None
404+
405+
init_score = jnp.zeros((b, t, k.shape[1]), dtype=jnp.float32)
406+
indexer_score, _ = jax.lax.scan(jax.checkpoint(scan_body_indexer), init_score, {"q": q_h, "w": w_h})
407+
indexer_score = indexer_score.astype(q.dtype)
408+
409+
else:
410+
# Aggregate head-wise logits: logits @ weights natively
411+
logits = jnp.einsum("bthd, bsd -> btsh", q, k, precision=self.config.matmul_precision)
412+
logits = jax.nn.relu(logits)
413+
indexer_score = jnp.einsum("btsh, bth -> bts", logits, weights, precision=self.config.matmul_precision)
376414

377415
internal_padding_mask = None
378416
if cached_s is not None:
@@ -1086,25 +1124,68 @@ def calculate_indexer_loss(
10861124
query = jax.lax.stop_gradient(query)
10871125
key = jax.lax.stop_gradient(key)
10881126

1089-
# Compute attention scores: [b, t, h, d] @ [b, s, h, d] -> [b, h, t, s]
1090-
attention_scores = jnp.einsum("bthd, bshd -> bhts", query, key, precision=self.config.matmul_precision)
1091-
1127+
# Ensure indexer_score updates identically in all branches
10921128
if sparse_loss:
1093-
# indexer_mask is already pre-filtered with the attention_mask if any
1094-
attention_scores = attention_scores + indexer_mask[:, None, :, :]
10951129
indexer_score = indexer_score + indexer_mask
1096-
elif attention_mask is not None:
1097-
# indexer_score already applies attention_mask; updating attention_scores only
1098-
attention_scores = attention_scores + attention_mask[:, None, :, :]
1099-
1100-
# Use float32 for softmax numerical stability.
1101-
attention_probs = jax.nn.softmax(attention_scores.astype(jnp.float32), axis=-1)
11021130
indexer_probs = jax.nn.softmax(indexer_score.astype(jnp.float32), axis=-1)
11031131

1104-
# Aggregate heads: [b, h, t, s] -> [b, t, s]
1105-
attention_probs = jnp.sum(attention_probs, axis=1)
1106-
# Force materialization and prevent fusion across this point to reuse the intermediate tensor
1107-
attention_probs = jax.lax.optimization_barrier(attention_probs)
1132+
batch, q_len, heads, dim = query.shape
1133+
1134+
# Chunk across the 'heads' dimension manually using jax.lax.scan
1135+
# Control the HBM footprint of QK tensor: [batch, q_len, s_len, heads]
1136+
# If set to 0, it falls back to native implementation.
1137+
head_chunk_size = getattr(self.config, "qk_head_chunk_size", 0)
1138+
if head_chunk_size > 0:
1139+
if head_chunk_size > heads or heads % head_chunk_size != 0:
1140+
raise ValueError(
1141+
f"qk_head_chunk_size ({head_chunk_size}) must be <= number of heads ({heads}) " f"and divide it evenly."
1142+
)
1143+
num_chunks = heads // head_chunk_size
1144+
1145+
# Transpose and reshape to put chunk dimension first for jax.lax.scan
1146+
# query: [b, t, h, d] -> [h, b, t, d] -> [num_chunks, head_chunk_size, b, t, d]
1147+
q_h = query.transpose(2, 0, 1, 3).reshape(num_chunks, head_chunk_size, batch, q_len, dim)
1148+
k_h = key.transpose(2, 0, 1, 3).reshape(num_chunks, head_chunk_size, batch, key.shape[1], dim)
1149+
1150+
def scan_body_heads(carry, xs):
1151+
q_c = xs["q"] # [h_chunk, b, t, d]
1152+
k_c = xs["k"] # [h_chunk, b, s, d]
1153+
1154+
# Directly use the chunked shapes in einsum to avoid transposes inside the loop
1155+
attn_chunk = jnp.einsum(
1156+
"hbtd, hbsd -> bhts",
1157+
q_c,
1158+
k_c,
1159+
precision=self.config.matmul_precision,
1160+
)
1161+
1162+
if sparse_loss:
1163+
attn_chunk = attn_chunk + indexer_mask[:, None, :, :]
1164+
elif attention_mask is not None:
1165+
attn_chunk = attn_chunk + attention_mask[:, None, :, :]
1166+
1167+
probs_chunk = jax.nn.softmax(attn_chunk.astype(jnp.float32), axis=-1)
1168+
probs_chunk_sum = jnp.sum(probs_chunk, axis=1) # [b, t, s]
1169+
1170+
return carry + probs_chunk_sum, None
1171+
1172+
init_probs = jnp.zeros((batch, q_len, key.shape[1]), dtype=jnp.float32)
1173+
attention_probs, _ = jax.lax.scan(jax.checkpoint(scan_body_heads), init_probs, {"q": q_h, "k": k_h})
1174+
1175+
else:
1176+
# Native implementation (default) if chunking is disabled
1177+
attention_scores = jnp.einsum(
1178+
"bthd, bshd -> bhts",
1179+
query,
1180+
key,
1181+
precision=self.config.matmul_precision,
1182+
)
1183+
if sparse_loss:
1184+
attention_scores = attention_scores + indexer_mask[:, None, :, :]
1185+
elif attention_mask is not None:
1186+
attention_scores = attention_scores + attention_mask[:, None, :, :]
1187+
attention_probs = jnp.sum(jax.nn.softmax(attention_scores.astype(jnp.float32), axis=-1), axis=1)
1188+
attention_probs = jax.lax.optimization_barrier(attention_probs)
11081189
# L1 normalize aggregated target distribution
11091190
attention_probs = attention_probs / (jnp.sum(attention_probs, axis=-1, keepdims=True) + EPS)
11101191

0 commit comments

Comments
 (0)