ollama/x/mlxrunner/batch/batch.go
Jesse Gross 2bbe2405fe mlxrunner: decouple models from attention cache storage layout
Models build their own attention masks and read K/V directly from
the cache's buffers, which ties them to the cache's storage layout.
That blocks multi-sequence batching — right-padded rows need a
query-padding mask composed onto every model — and rules out
variants like paged attention where K/V isn't one contiguous tensor.

Caches now hand back a per-layer KVHistory holding post-update K, V,
and a MaskApplier that merges the cache's storage restrictions into
the model's logical mask. Models describe their mask in logical
terms; SDPA composes model, padding, and applier contributions and
dispatches to the kernel's causal or no-mask fast path when it can.
KVHistory still exposes K, V, and the composed mask for manual
attention paths (e.g. CUDA prefill at head_dim > 128).

Performance for single-sequence inference is unchanged.
2026-04-27 20:04:46 -07:00

42 lines
1.1 KiB
Go

package batch
import "github.com/ollama/ollama/x/mlxrunner/mlx"
// Batch is the per-forward-pass input handed to a model.
type Batch struct {
// InputIDs is the input token IDs for this forward pass, shape (B, L).
InputIDs *mlx.Array
// SeqOffsets gives each row's current position within its sequence —
// where the chunk in InputIDs starts. Length equals the batch dimension
// of InputIDs.
SeqOffsets []int32
// SeqQueryLens is each row's real query length in this forward. Values
// less than L mean the row's tail is padding that must be masked out.
// Length equals the batch dimension of InputIDs.
SeqQueryLens []int32
// Memo is per-forward memoization used to cache results, such as masks,
// which are often the same across layers.
Memo Memo
}
type Memo struct {
entries map[any]any
}
// Get returns the memoized value for key and true if present, or nil
// and false otherwise.
func (m *Memo) Get(key any) (any, bool) {
v, ok := m.entries[key]
return v, ok
}
// Put stores value under key, allocating on first use.
func (m *Memo) Put(key, value any) {
if m.entries == nil {
m.entries = map[any]any{}
}
m.entries[key] = value
}