Skip to content

Commit 308648c

Browse files
kernelpoolawni
authored andcommitted
Fix sparse token selection in deepseek v3.2 (#531)
* Fix sparse token selection in deepseek v3.2 * Fix 4D mask input handling and remove unnecessary ones array
1 parent da16ea9 commit 308648c

1 file changed

Lines changed: 14 additions & 19 deletions

File tree

mlx_lm/models/deepseek_v32.py

Lines changed: 14 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -101,10 +101,12 @@ def __call__(
101101
k = cache.update_and_fetch(k)
102102
if k.shape[2] <= self.index_topk:
103103
return None
104+
scores = q @ k.swapaxes(-1, -2)
105+
scores = mx.maximum(scores, 0)
104106
weights = self.weights_proj(x) * (self.n_heads**-0.5)
105107
weights = (weights * self.softmax_scale).swapaxes(-1, -2)[..., None]
106-
q_scaled = q * weights
107-
scores = (q * weights) @ k.swapaxes(-1, -2)
108+
scores = scores * weights
109+
scores = scores.sum(axis=1)
108110
if mask is not None:
109111
scores = mx.where(mask, scores, -float("inf"))
110112
return mx.argpartition(scores, kth=-self.index_topk, axis=-1)[
@@ -219,24 +221,17 @@ def __call__(
219221
queries = mx.concatenate([q_nope, q_pe], axis=-1)
220222
topk_indices = self.indexer(x, qr, mask, cache=cache[1])
221223
if topk_indices is not None:
222-
repeats = self.num_heads // self.config.index_n_heads
223-
if L == 1:
224-
topk_indices = mx.repeat(topk_indices, repeats, axis=1).squeeze(-2)[
225-
..., None
226-
]
227-
keys = mx.take_along_axis(keys, topk_indices, axis=-2)
228-
values = mx.take_along_axis(values, topk_indices, axis=-2)
229-
else:
230-
topk_mask = mx.zeros(
231-
(B, self.config.index_n_heads, *mask.shape[-2:]), mx.bool_
232-
)
233-
topk_mask = mx.put_along_axis(
234-
topk_mask, topk_indices, mx.array(True), axis=-1
235-
)
236-
mask = mask & topk_mask
237-
mask = mx.repeat(mask, repeats, axis=1)
224+
k_seq = keys.shape[2]
225+
sparse_mask = mx.zeros((B, L, k_seq), dtype=mx.bool_)
226+
sparse_mask = mx.put_along_axis(
227+
sparse_mask, topk_indices, mx.array(True), axis=-1
228+
)
229+
sparse_mask = sparse_mask[:, None, :, :]
230+
if mask is not None:
231+
sparse_mask = sparse_mask & mask
232+
mask = sparse_mask
238233
output = scaled_dot_product_attention(
239-
queries, keys, values, cache=cache, scale=self.scale, mask=mask
234+
queries, keys, values, cache=cache[0], scale=self.scale, mask=mask
240235
)
241236
output = output.transpose(0, 2, 1, 3).reshape(B, L, -1)
242237
return self.o_proj(output)

0 commit comments

Comments
 (0)