[Apple Silicon][MLX] Cache seq_lens-derived tensors in BatchedDecodeContext (#23470)

Signed-off-by: Xiaodong Ye <yeahdongcn@gmail.com>
This commit is contained in:
R0CKSTAR
2026-04-23 18:12:26 -07:00
committed by GitHub
parent 74c2e5bacd
commit 87e50f20f6
@@ -3,7 +3,7 @@
from __future__ import annotations from __future__ import annotations
import threading import threading
from dataclasses import dataclass from dataclasses import dataclass, field
from typing import Any, Optional from typing import Any, Optional
import mlx.core as mx import mlx.core as mx
@@ -23,6 +23,24 @@ class BatchedDecodeContext:
# layer_caches[layer_idx][req_idx] = ContiguousKVCache # layer_caches[layer_idx][req_idx] = ContiguousKVCache
layer_caches: list[list[ContiguousKVCache]] layer_caches: list[list[ContiguousKVCache]]
# Derived tensors/metadata, shared across all layers in one forward pass.
offsets: mx.array = field(init=False)
max_len: int = field(init=False)
valid_lens: mx.array = field(init=False)
needs_padding: bool = field(init=False)
pad_sizes: list[int] = field(init=False)
positions: Optional[mx.array] = field(init=False)
def __post_init__(self) -> None:
seq_lens = self.seq_lens
max_seq_len = max(seq_lens)
self.offsets = mx.array(seq_lens, dtype=mx.int32)
self.max_len = max_seq_len + 1
self.valid_lens = self.offsets + 1
self.needs_padding = min(seq_lens) < max_seq_len
self.pad_sizes = [max_seq_len - s for s in seq_lens]
self.positions = mx.arange(self.max_len) if self.needs_padding else None
def set_context(ctx: Optional[BatchedDecodeContext]) -> None: def set_context(ctx: Optional[BatchedDecodeContext]) -> None:
_thread_local.batched_ctx = ctx _thread_local.batched_ctx = ctx
@@ -78,12 +96,13 @@ class MLXAttentionWrapper(nn.Module):
values = values.transpose(0, 2, 1, 3) values = values.transpose(0, 2, 1, 3)
# Vectorized RoPE with per-batch offsets # Vectorized RoPE with per-batch offsets
offsets = mx.array(ctx.seq_lens, dtype=mx.int32) offsets = ctx.offsets
queries = inner.rope(queries, offset=offsets) queries = inner.rope(queries, offset=offsets)
keys = inner.rope(keys, offset=offsets) keys = inner.rope(keys, offset=offsets)
layer_caches = ctx.layer_caches[layer_idx] layer_caches = ctx.layer_caches[layer_idx]
max_len = max(ctx.seq_lens) + 1 max_len = ctx.max_len
pad_sizes = ctx.pad_sizes
# TODO: replace per-request loop with native batched/ragged # TODO: replace per-request loop with native batched/ragged
# attention once mx.fast.scaled_dot_product_attention supports # attention once mx.fast.scaled_dot_product_attention supports
@@ -95,10 +114,9 @@ class MLXAttentionWrapper(nn.Module):
layer_caches[i].write_token(keys[i : i + 1], values[i : i + 1]) layer_caches[i].write_token(keys[i : i + 1], values[i : i + 1])
k_all, v_all = layer_caches[i].get_kv() k_all, v_all = layer_caches[i].get_kv()
curr_len = layer_caches[i].offset
if curr_len < max_len: pad = pad_sizes[i]
pad = max_len - curr_len if pad > 0:
k_pad = mx.zeros( k_pad = mx.zeros(
(1, inner.n_kv_heads, pad, head_dim), dtype=k_all.dtype (1, inner.n_kv_heads, pad, head_dim), dtype=k_all.dtype
) )
@@ -115,11 +133,8 @@ class MLXAttentionWrapper(nn.Module):
values_b = mx.concatenate(all_v, axis=0) values_b = mx.concatenate(all_v, axis=0)
attn_mask = None attn_mask = None
seq_lens_plus1 = [s + 1 for s in ctx.seq_lens] if ctx.needs_padding:
if min(seq_lens_plus1) < max_len: mask_bool = ctx.positions[None, :] >= ctx.valid_lens[:, None]
positions = mx.arange(max_len)
valid_lens = mx.array(seq_lens_plus1, dtype=mx.int32)
mask_bool = positions[None, :] >= valid_lens[:, None]
attn_mask = mx.where( attn_mask = mx.where(
mask_bool[:, None, None, :], mask_bool[:, None, None, :],
mx.array(mx.finfo(queries.dtype).min, dtype=queries.dtype), mx.array(mx.finfo(queries.dtype).min, dtype=queries.dtype),