Skip to content

Commit 53dffbd

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 53dffbd

5 files changed

Lines changed: 295 additions & 26 deletions

File tree

patch_mla.py

Lines changed: 163 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,163 @@
1+
import sys
2+
3+
with open("src/maxtext/layers/attention_mla.py", "r") as f:
4+
code = f.read()
5+
6+
target1 = """ # Aggregate head-wise logits: logits @ weights
7+
# [b, t, h, d] @ [b, s, d] -> [b, t, s, h]
8+
logits = jnp.einsum("bthd, bsd -> btsh", q, k, precision=self.config.matmul_precision)
9+
logits = jax.nn.relu(logits)
10+
# [b, t, s, h] @ [b, t, h] -> [b, t, s]
11+
indexer_score = jnp.einsum("btsh, bth -> bts", logits, weights, precision=self.config.matmul_precision)"""
12+
13+
replace1 = """ # Compute Index Scores safely by chunking the 'heads' dimension to prevent OOM
14+
# The naive evaluation materializes [b, t, s, h] which explodes HBM.
15+
# We use jax.lax.scan to compute the score iteratively over head chunks.
16+
b, t, h, d = q.shape
17+
# Control the HBM footprint of QK tensor: [batch, q_len, s_len, heads] or [batch, heads, q_len, s_len]
18+
# If set to 0, it falls back to native materialization.
19+
head_chunk_size = getattr(self.config, "qk_head_chunk_size", 0)
20+
if head_chunk_size > 0:
21+
if head_chunk_size > h or h % head_chunk_size != 0:
22+
raise ValueError(
23+
f"qk_head_chunk_size ({head_chunk_size}) must be <= number of heads ({h}) "
24+
f"and divide it evenly."
25+
)
26+
num_chunks = h // head_chunk_size
27+
# q: [b, t, h, d] -> [h, b, t, d] -> [num_chunks, head_chunk_size, b, t, d]
28+
q_h = q.transpose(2, 0, 1, 3).reshape(num_chunks, head_chunk_size, b, t, d)
29+
# weights: [b, t, h] -> [h, b, t] -> [num_chunks, head_chunk_size, b, t]
30+
w_h = weights.transpose(2, 0, 1).reshape(num_chunks, head_chunk_size, b, t)
31+
32+
def scan_body_indexer(carry, xs):
33+
q_c = xs["q"] # [h_chunk, b, t, d]
34+
w_c = xs["w"] # [h_chunk, b, t]
35+
36+
# Transpose back for einsum:
37+
q_c = q_c.transpose(1, 2, 0, 3) # [b, t, h_chunk, d]
38+
w_c = w_c.transpose(1, 2, 0) # [b, t, h_chunk]
39+
40+
logits = jnp.einsum("bthd, bsd -> btsh", q_c, k, precision=self.config.matmul_precision)
41+
logits = jax.nn.relu(logits)
42+
43+
# Upcast the chunk to fp32 directly inside the block so we precisely match MXU
44+
# accumulator numerical stability without bfloat16 degradation over the loop boundary
45+
score_chunk = jnp.einsum(
46+
"btsh, bth -> bts",
47+
logits,
48+
w_c,
49+
precision=self.config.matmul_precision,
50+
)
51+
return carry + score_chunk.astype(jnp.float32), None
52+
53+
init_score = jnp.zeros((b, t, k.shape[1]), dtype=jnp.float32)
54+
indexer_score, _ = jax.lax.scan(jax.checkpoint(scan_body_indexer), init_score, {"q": q_h, "w": w_h})
55+
indexer_score = indexer_score.astype(q.dtype)
56+
57+
else:
58+
# Aggregate head-wise logits: logits @ weights natively
59+
logits = jnp.einsum("bthd, bsd -> btsh", q, k, precision=self.config.matmul_precision)
60+
logits = jax.nn.relu(logits)
61+
indexer_score = jnp.einsum("btsh, bth -> bts", logits, weights, precision=self.config.matmul_precision)"""
62+
63+
target2 = """ # Compute attention scores: [b, t, h, d] @ [b, s, h, d] -> [b, h, t, s]
64+
attention_scores = jnp.einsum("bthd, bshd -> bhts", query, key, precision=self.config.matmul_precision)
65+
66+
if sparse_loss:
67+
# indexer_mask is already pre-filtered with the attention_mask if any
68+
attention_scores = attention_scores + indexer_mask[:, None, :, :]
69+
indexer_score = indexer_score + indexer_mask
70+
elif attention_mask is not None:
71+
# indexer_score already applies attention_mask; updating attention_scores only
72+
attention_scores = attention_scores + attention_mask[:, None, :, :]
73+
74+
# Use float32 for softmax numerical stability.
75+
attention_probs = jax.nn.softmax(attention_scores.astype(jnp.float32), axis=-1)
76+
indexer_probs = jax.nn.softmax(indexer_score.astype(jnp.float32), axis=-1)
77+
78+
# Aggregate heads: [b, h, t, s] -> [b, t, s]
79+
attention_probs = jnp.sum(attention_probs, axis=1)
80+
# Force materialization and prevent fusion across this point to reuse the intermediate tensor
81+
attention_probs = jax.lax.optimization_barrier(attention_probs)"""
82+
83+
replace2 = """ # Ensure indexer_score updates identically in all branches
84+
if sparse_loss:
85+
indexer_score = indexer_score + indexer_mask
86+
indexer_probs = jax.nn.softmax(indexer_score.astype(jnp.float32), axis=-1)
87+
88+
batch, q_len, heads, dim = query.shape
89+
90+
# Chunk across the 'heads' dimension manually using jax.lax.scan
91+
# Control the HBM footprint of QK tensor: [batch, q_len, s_len, heads]
92+
# If set to 0, it falls back to native implementation.
93+
head_chunk_size = getattr(self.config, "qk_head_chunk_size", 0)
94+
if head_chunk_size > 0:
95+
if head_chunk_size > heads or heads % head_chunk_size != 0:
96+
raise ValueError(
97+
f"qk_head_chunk_size ({head_chunk_size}) must be <= number of heads ({heads}) "
98+
f"and divide it evenly."
99+
)
100+
num_chunks = heads // head_chunk_size
101+
102+
# Transpose and reshape to put chunk dimension first for jax.lax.scan
103+
# query: [b, t, h, d] -> [h, b, t, d] -> [num_chunks, head_chunk_size, b, t, d]
104+
q_h = query.transpose(2, 0, 1, 3).reshape(num_chunks, head_chunk_size, batch, q_len, dim)
105+
k_h = key.transpose(2, 0, 1, 3).reshape(num_chunks, head_chunk_size, batch, key.shape[1], dim)
106+
107+
def scan_body_heads(carry, xs):
108+
q_c = xs["q"] # [h_chunk, b, t, d]
109+
k_c = xs["k"] # [h_chunk, b, s, d]
110+
111+
# Transpose back: [h_chunk, b, t, d] -> [b, t, h_chunk, d]
112+
q_c = q_c.transpose(1, 2, 0, 3)
113+
k_c = k_c.transpose(1, 2, 0, 3)
114+
115+
attn_chunk = jnp.einsum(
116+
"bthd, bshd -> bhts",
117+
q_c,
118+
k_c,
119+
precision=self.config.matmul_precision,
120+
)
121+
122+
if sparse_loss:
123+
attn_chunk = attn_chunk + indexer_mask[:, None, :, :]
124+
elif attention_mask is not None:
125+
attn_chunk = attn_chunk + attention_mask[:, None, :, :]
126+
127+
probs_chunk = jax.nn.softmax(attn_chunk.astype(jnp.float32), axis=-1)
128+
probs_chunk_sum = jnp.sum(probs_chunk, axis=1) # [b, t, s]
129+
130+
return carry + probs_chunk_sum, None
131+
132+
init_probs = jnp.zeros((batch, q_len, key.shape[1]), dtype=jnp.float32)
133+
attention_probs, _ = jax.lax.scan(jax.checkpoint(scan_body_heads), init_probs, {"q": q_h, "k": k_h})
134+
135+
else:
136+
# Natively evaluate if chunking is disabled or head sizing does not allow clean chunks
137+
attention_scores = jnp.einsum(
138+
"bthd, bshd -> bhts",
139+
query,
140+
key,
141+
precision=self.config.matmul_precision,
142+
)
143+
if sparse_loss:
144+
attention_scores = attention_scores + indexer_mask[:, None, :, :]
145+
elif attention_mask is not None:
146+
attention_scores = attention_scores + attention_mask[:, None, :, :]
147+
attention_probs = jnp.sum(jax.nn.softmax(attention_scores.astype(jnp.float32), axis=-1), axis=1)
148+
attention_probs = jax.lax.optimization_barrier(attention_probs)"""
149+
150+
if target1 not in code:
151+
print("Could not find target1")
152+
sys.exit(1)
153+
if target2 not in code:
154+
print("Could not find target2")
155+
sys.exit(1)
156+
157+
code = code.replace(target1, replace1)
158+
code = code.replace(target2, replace2)
159+
160+
with open("src/maxtext/layers/attention_mla.py", "w") as f:
161+
f.write(code)
162+
163+
print("Patch fully completed.")

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: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -637,6 +637,16 @@ 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+
ge=0,
643+
description=(
644+
"Chunk size over heads dimension for QK attention dot product. "
645+
"Default is 0 (no chunking). Must divide both the attention query "
646+
"heads (num_query_heads) and indexer heads (indexer_n_heads) evenly. "
647+
"This reduces memory footprint at the cost of time, especially helpful in long context."
648+
),
649+
)
640650

641651

642652
class MoBa(BaseModel):

src/maxtext/layers/attention_mla.py

Lines changed: 104 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 indexer heads ({h}) " "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,69 @@ 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 attention query heads ({heads}) "
1142+
"and divide it evenly."
1143+
)
1144+
num_chunks = heads // head_chunk_size
1145+
1146+
# Transpose and reshape to put chunk dimension first for jax.lax.scan
1147+
# query: [b, t, h, d] -> [h, b, t, d] -> [num_chunks, head_chunk_size, b, t, d]
1148+
q_h = query.transpose(2, 0, 1, 3).reshape(num_chunks, head_chunk_size, batch, q_len, dim)
1149+
k_h = key.transpose(2, 0, 1, 3).reshape(num_chunks, head_chunk_size, batch, key.shape[1], dim)
1150+
1151+
def scan_body_heads(carry, xs):
1152+
q_c = xs["q"] # [h_chunk, b, t, d]
1153+
k_c = xs["k"] # [h_chunk, b, s, d]
1154+
1155+
# Directly use the chunked shapes in einsum to avoid transposes inside the loop
1156+
attn_chunk = jnp.einsum(
1157+
"hbtd, hbsd -> bhts",
1158+
q_c,
1159+
k_c,
1160+
precision=self.config.matmul_precision,
1161+
)
1162+
1163+
if sparse_loss:
1164+
attn_chunk = attn_chunk + indexer_mask[:, None, :, :]
1165+
elif attention_mask is not None:
1166+
attn_chunk = attn_chunk + attention_mask[:, None, :, :]
1167+
1168+
probs_chunk = jax.nn.softmax(attn_chunk.astype(jnp.float32), axis=-1)
1169+
probs_chunk_sum = jnp.sum(probs_chunk, axis=1) # [b, t, s]
1170+
1171+
return carry + probs_chunk_sum, None
1172+
1173+
init_probs = jnp.zeros((batch, q_len, key.shape[1]), dtype=jnp.float32)
1174+
attention_probs, _ = jax.lax.scan(scan_body_heads, init_probs, {"q": q_h, "k": k_h})
1175+
1176+
else:
1177+
# Native implementation (default) if chunking is disabled
1178+
attention_scores = jnp.einsum(
1179+
"bthd, bshd -> bhts",
1180+
query,
1181+
key,
1182+
precision=self.config.matmul_precision,
1183+
)
1184+
if sparse_loss:
1185+
attention_scores = attention_scores + indexer_mask[:, None, :, :]
1186+
elif attention_mask is not None:
1187+
attention_scores = attention_scores + attention_mask[:, None, :, :]
1188+
attention_probs = jnp.sum(jax.nn.softmax(attention_scores.astype(jnp.float32), axis=-1), axis=1)
1189+
attention_probs = jax.lax.optimization_barrier(attention_probs)
11081190
# L1 normalize aggregated target distribution
11091191
attention_probs = attention_probs / (jnp.sum(attention_probs, axis=-1, keepdims=True) + EPS)
11101192

0 commit comments

Comments
 (0)