mlxrunner: support draft heads that maintain draft caches

Generalize the draft path so a head that maintains a KV cache (EAGLE-style)
and Gemma's read-only single-position assistant both fit one drafter
interface with no per-model branches, and make the committed stream the
drafter's maintenance mechanism — every committed run is reported, the
drafter pairs each draft slot with its look-ahead token and flushes completed
pairs to the draft caches. The draft KV thus stays prefix-cached alongside
the target in every session, drafting or not.
This commit is contained in:
Jesse Gross
2026-06-09 17:10:40 -07:00
parent 340c51bbb7
commit 28fbbb06d5
8 changed files with 791 additions and 162 deletions

View File

@@ -17,6 +17,10 @@ type Batch struct {
// Length equals the batch dimension of InputIDs. // Length equals the batch dimension of InputIDs.
SeqQueryLens []int32 SeqQueryLens []int32
// Hidden is the target hidden state a draft model fuses with its input
// embedding for this step. It is nil for ordinary forward passes.
Hidden *mlx.Array
// Memo is per-forward memoization used to cache results, such as masks, // Memo is per-forward memoization used to cache results, such as masks,
// which are often the same across layers. // which are often the same across layers.
Memo Memo Memo Memo

View File

@@ -27,8 +27,21 @@ type Model interface {
LoadWeights(tensors map[string]*mlx.Array) error LoadWeights(tensors map[string]*mlx.Array) error
} }
// DraftModel is an auxiliary model stored alongside a target model. // DraftModel is an auxiliary model alongside a target that proposes speculative
// tokens.
type DraftModel interface { type DraftModel interface {
// Draft fuses b.Hidden (the target hidden state) into its own forward and
// returns the head's hidden plus the projected hidden seeding the next step.
Draft(b *batch.Batch, caches []cache.Cache) (hidden, projected *mlx.Array)
// Unembed projects a hidden state to vocabulary logits.
Unembed(x *mlx.Array) *mlx.Array
// DraftCaches selects the draft model's own KV caches from the full
// per-request slice — any subset, or nil when the draft keeps no KV.
DraftCaches(caches []cache.Cache) []cache.Cache
// LoadWeights assigns manifest tensors to the draft head's fields.
LoadWeights(tensors map[string]*mlx.Array) error LoadWeights(tensors map[string]*mlx.Array) error
} }
@@ -46,15 +59,10 @@ type MTPDefaultsProvider interface {
MTPDraftDefaults(sample bool) MTPDefaults MTPDraftDefaults(sample bool) MTPDefaults
} }
// MTPDraftModel is a draft model capable of Gemma-style multi-token // SelfDraft is implemented by models whose draft head ships inline with the
// prediction from target token embeddings, target hidden states, and target KV. // target weights; it returns the head, or nil when the checkpoint shipped none.
type MTPDraftModel interface { type SelfDraft interface {
Draft(inputEmbeds *mlx.Array, position int32, caches []cache.Cache) (logits, hidden *mlx.Array) SelfDraft() DraftModel
}
// MTPEmbeddingModel exposes the target token embedding path used by MTP drafts.
type MTPEmbeddingModel interface {
TokenEmbeddings(inputIDs *mlx.Array) *mlx.Array
} }
var ( var (

View File

@@ -1,7 +1,9 @@
package mlxrunner package mlxrunner
import ( import (
"github.com/ollama/ollama/x/mlxrunner/cache" "fmt"
"github.com/ollama/ollama/x/mlxrunner/batch"
"github.com/ollama/ollama/x/mlxrunner/mlx" "github.com/ollama/ollama/x/mlxrunner/mlx"
"github.com/ollama/ollama/x/mlxrunner/model/base" "github.com/ollama/ollama/x/mlxrunner/model/base"
sampler "github.com/ollama/ollama/x/mlxrunner/sample" sampler "github.com/ollama/ollama/x/mlxrunner/sample"
@@ -12,6 +14,11 @@ const (
mtpDefaultMaxDraftTokens = 16 mtpDefaultMaxDraftTokens = 16
) )
// mtpPendingFlushTokens caps how many committed look-ahead tokens wait in the
// pending buffer before a batched flush, bounding the pinned hidden states
// regardless of what else triggers a flush.
const mtpPendingFlushTokens = 32
func (r *Runner) mtpDefaults(sample bool) base.MTPDefaults { func (r *Runner) mtpDefaults(sample bool) base.MTPDefaults {
defaults := base.MTPDefaults{ defaults := base.MTPDefaults{
InitialDraftTokens: mtpDefaultInitialDraftTokens, InitialDraftTokens: mtpDefaultInitialDraftTokens,
@@ -30,85 +37,231 @@ func (r *Runner) mtpDefaults(sample bool) base.MTPDefaults {
return defaults return defaults
} }
// mtpDrafter drafts with a model's multi-token-prediction head: a small // mtpDrafter drafts with a model's multi-token-prediction head, fed through
// head trained to continue the target's hidden states, fed through the // the committed-stream reports. The draft KV pairs each slot S with the
// committed-stream reports. It retains the hidden at the last committed // look-ahead token at S+1 fused with the target hidden at S, so a pair
// slot and tracks the committed frontier, which anchors its drafting. // completes only when the next token arrives.
type mtpDrafter struct { type mtpDrafter struct {
spec *speculation spec *speculation
draft base.MTPDraftModel
target base.MTPEmbeddingModel
caches []cache.Cache
// frontier is the slot after the last committed token; the hidden // frontier is the slot after the last reported token; frontierHidden is
// reported for slot frontier-1 is retained (pinned) as the fusion // the pinned target hidden at frontier-1, fused into the next pair.
// input for the next proposal chain.
frontier int frontier int
frontierHidden *mlx.Array frontierHidden *mlx.Array
// committedDraftOffset is the slot after the last pair written to the
// draft caches; later pairs wait pinned in the pending lists until
// flushed. pendingCount is the look-ahead tokens those lists hold, summed
// across the buffered runs.
committedDraftOffset int
pendingTokens []*mlx.Array
pendingHiddens []*mlx.Array
pendingCount int
// heldHidden is the frontier row's pre-unembed hidden and heldProjected
// its fusion hidden, carried from the last flush so the first proposal
// reuses them without a head call.
heldHidden *mlx.Array
heldProjected *mlx.Array
} }
// newMTPDrafter returns the MTP drafter for this request's caches, or nil // newMTPDrafter returns the MTP drafter cursor for this request, syncing its
// when the model carries no MTP head. // pairing frontier to the draft caches' restored offset.
func newMTPDrafter(s *speculation, caches []cache.Cache) *mtpDrafter { func newMTPDrafter(s *speculation) *mtpDrafter {
draft, ok := s.draft.(base.MTPDraftModel) d := &mtpDrafter{spec: s}
if !ok { if len(s.draftKV) > 0 {
return nil // A restored prefix arrives with the draft caches already written;
// pairing resumes from their absolute offset.
d.committedDraftOffset = s.draftKV[0].Offset()
d.frontier = d.committedDraftOffset
} }
target, ok := s.r.Model.(base.MTPEmbeddingModel) return d
if !ok {
return nil
}
return &mtpDrafter{spec: s, draft: draft, target: target, caches: caches}
} }
func (d *mtpDrafter) committed(tokens, hiddens *mlx.Array, position int) { func (d *mtpDrafter) committed(tokens, hiddens *mlx.Array, position int) {
d.frontier = position + tokens.Dim(1) n := tokens.Dim(1)
h := lastHiddenRow(hiddens) if len(d.spec.draftKV) > 0 {
mlx.Pin(h) // The pair at slot S fuses token[S+1] with hidden[S], so a run pairs its
if d.frontierHidden != nil { // tokens with its own hiddens shifted one slot back: the first writable
mlx.Unpin(d.frontierHidden) // token takes the carried frontier hidden, each later token the row
// before it. Leading tokens whose slot is already buffered or written
// through (a proposal consumed the run's first token, or a restored
// prefix sits at the run start) are skipped; a slot below the frontier
// is a gap bug.
start := d.committedDraftOffset + d.pendingCount - position + 1
if start < 0 {
panic(fmt.Sprintf("mtp: committed run at %d leaves a pair gap at %d", position, d.committedDraftOffset+d.pendingCount))
}
if start < n {
ids := tokens.Slice(mlx.Slice(), mlx.Slice(start, n))
var h *mlx.Array
if start == 0 {
h = mlx.Concatenate([]*mlx.Array{d.frontierHidden, hiddens.Slice(mlx.Slice(), mlx.Slice(0, n-1), mlx.Slice())}, 1)
} else {
h = hiddens.Slice(mlx.Slice(), mlx.Slice(start-1, n-1), mlx.Slice())
}
d.queueCacheWrites(ids, h)
}
} }
d.frontierHidden = h
d.frontier = position + n
d.setFrontierHidden(lastHiddenRow(hiddens))
}
// finish settles the drafter when generation ends: current completes the
// frontier pair, leveling the draft caches with the target's resting offset.
//
// TODO: leveling the draft to the target writes a boundary entry whose
// look-ahead token is outside the stored prefix (here, the never-committed
// current). When a later request restores this prefix and diverges at the
// boundary, that entry is stale and lowers draft acceptance. EAGLE keeps the
// draft one slot behind the target so the unconfirmed boundary entry is never
// written (its bigram partner token[S+1] does not exist yet); we should do the
// same rather than level here. Regenerating hidden[S] to re-pair the boundary
// on the next request is a separate, re-prefill-bound concern for recurrent
// targets.
func (d *mtpDrafter) finish(current *mlx.Array) {
if len(d.spec.draftKV) == 0 {
return
}
d.settle(current)
}
// settle completes any open frontier pair with current, then flushes.
func (d *mtpDrafter) settle(current *mlx.Array) {
if d.frontierHidden != nil && d.frontier-1 == d.committedDraftOffset+d.pendingCount {
d.queueCacheWrites(current.ExpandDims(-1), d.frontierHidden)
}
d.flush()
} }
func (d *mtpDrafter) close() { func (d *mtpDrafter) close() {
if d.frontierHidden != nil { d.flush()
mlx.Unpin(d.frontierHidden) d.setFrontierHidden(nil)
d.frontierHidden = nil d.setHeld(nil, nil)
}
// queueCacheWrites buffers completed draft-cache writes — look-ahead tokens
// fused with their target hiddens — flushing once the buffer reaches the token
// cap so the pinned hiddens stay bounded. flush coalesces the buffered writes
// into one head forward, so a contiguous run lands in a single draft-cache extend.
func (d *mtpDrafter) queueCacheWrites(tokens, hiddens *mlx.Array) {
mlx.Pin(tokens, hiddens)
d.pendingTokens = append(d.pendingTokens, tokens)
d.pendingHiddens = append(d.pendingHiddens, hiddens)
d.pendingCount += tokens.Dim(1)
if d.pendingCount >= mtpPendingFlushTokens {
d.flush()
} }
} }
// propose drafts a token chain Gemma-style ("single-position"): the head is // flush writes the pending pairs to the draft caches in one head forward,
// trained to draft every speculative token as if it sat at the last // dropping speculative entries past the committed range first and holding
// committed slot, re-attending the target caches read-only. Each step fuses // the last row's logits and projected hidden for the next proposal chain.
// the previous token's target embedding with the hidden chain — the func (d *mtpDrafter) flush() {
// reported hidden first, then the head's own projections — and the if len(d.pendingTokens) == 0 {
// RoPE/cache anchor stays at the last committed slot while the proposed return
// tokens advance. }
for _, c := range d.spec.draftKV {
if c.Offset() > d.committedDraftOffset {
if !c.Restore(nil, d.committedDraftOffset) {
panic(fmt.Sprintf("mtp: draft cache rewind to %d failed", d.committedDraftOffset))
}
}
}
ids := mlx.Concatenate(d.pendingTokens, 1)
hiddens := mlx.Concatenate(d.pendingHiddens, 1)
hidden, projected := d.spec.draft.Draft(&batch.Batch{
InputIDs: ids,
SeqOffsets: []int32{int32(d.committedDraftOffset)},
SeqQueryLens: []int32{int32(ids.Dim(1))},
Hidden: hiddens,
}, d.spec.caches)
d.setHeld(lastHiddenRow(hidden), lastHiddenRow(projected))
d.committedDraftOffset += ids.Dim(1)
// Force the draft writes: a session that never drafts would otherwise
// leave the flush chain unevaluated, pinning every hidden until close.
state := make([]*mlx.Array, 0, 2*len(d.spec.draftKV))
for _, c := range d.spec.draftKV {
state = append(state, c.State()...)
}
mlx.AsyncEval(state...)
mlx.Unpin(d.pendingTokens...)
mlx.Unpin(d.pendingHiddens...)
d.pendingTokens, d.pendingHiddens = nil, nil
d.pendingCount = 0
}
func (d *mtpDrafter) setFrontierHidden(h *mlx.Array) {
mlx.Pin(h)
mlx.Unpin(d.frontierHidden)
d.frontierHidden = h
}
// setHeld replaces the held flush outputs, pinned until the next flush or close.
func (d *mtpDrafter) setHeld(hidden, projected *mlx.Array) {
mlx.Pin(hidden, projected)
mlx.Unpin(d.heldHidden, d.heldProjected)
d.heldHidden, d.heldProjected = hidden, projected
}
// propose drafts a token chain after the not-yet-validated current token.
// A head with draft caches settles the frontier pair first, so its first step
// reuses the held frontier row with no head call; a cacheless head re-attends
// the target caches read-only, anchored at the last committed slot.
func (d *mtpDrafter) propose(current *mlx.Array, maxTokens int) *draftCandidates { func (d *mtpDrafter) propose(current *mlx.Array, maxTokens int) *draftCandidates {
if maxTokens <= 0 || d.frontierHidden == nil { if maxTokens <= 0 || d.frontierHidden == nil {
return nil return nil
} }
r := d.spec.r r := d.spec.r
anchor := int32(d.frontier - 1) if len(d.spec.draftKV) > 0 {
d.settle(current)
if d.heldHidden == nil {
return nil
}
}
lastToken := current.ExpandDims(-1) lastToken := current.ExpandDims(-1)
lastHidden := d.frontierHidden lastHidden := d.frontierHidden
draftTokens := make([]*mlx.Array, 0, maxTokens)
draftDists := make([]sampler.Distribution, 0, maxTokens) draftDists := make([]sampler.Distribution, 0, maxTokens)
var prefix *mlx.Array var prefix *mlx.Array
for range maxTokens { for i := range maxTokens {
tokenEmbedding := d.target.TokenEmbeddings(lastToken) var hidden, projected *mlx.Array
inputs := tokenEmbedding.Concatenate(-1, lastHidden) if i == 0 && len(d.spec.draftKV) > 0 {
logits, projected := d.draft.Draft(inputs, anchor, d.caches) // The settle flush already produced the frontier row; reuse it
stepLogits := lastLogits(logits) // instead of re-running the head.
hidden, projected = d.heldHidden, d.heldProjected
} else {
// A head with draft caches writes each draft token to the next
// draft-cache slot, advancing one per step from the last committed
// slot (the held i==0 step stands in for that slot). A cacheless
// head stays at the last committed slot every step, re-attending
// the committed prefix read-only ("single-position").
pos := d.frontier - 1
if len(d.spec.draftKV) > 0 {
pos = d.frontier - 1 + i
}
hidden, projected = d.spec.draft.Draft(&batch.Batch{
InputIDs: lastToken,
SeqOffsets: []int32{int32(pos)},
SeqQueryLens: []int32{1},
Hidden: lastHidden,
}, d.spec.caches)
}
// Unembed only the row being sampled, never the batch.
stepLogits := d.spec.draft.Unembed(hidden).Squeeze(1)
lastHidden = projected
// The chain's earlier drafts ride along as the row's history, so
// penalties shape proposals the same way they shape validation.
dist := r.Sampler.Distribution(pipelineSlot, stepLogits, prefix) dist := r.Sampler.Distribution(pipelineSlot, stepLogits, prefix)
nextToken := r.Sampler.SampleDistribution(pipelineSlot, dist) nextToken := r.Sampler.SampleDistribution(pipelineSlot, dist)
lastToken = nextToken.ExpandDims(-1) lastToken = nextToken.ExpandDims(-1)
lastHidden = projected
draftTokens = append(draftTokens, lastToken)
draftDists = append(draftDists, dist) draftDists = append(draftDists, dist)
if prefix == nil { if prefix == nil {
prefix = lastToken prefix = lastToken
@@ -117,7 +270,7 @@ func (d *mtpDrafter) propose(current *mlx.Array, maxTokens int) *draftCandidates
} }
} }
return &draftCandidates{ return &draftCandidates{
tokens: mlx.Concatenate(draftTokens, 1), tokens: prefix,
dist: sampler.ConcatenateDistributions(draftDists), dist: sampler.ConcatenateDistributions(draftDists),
} }
} }

View File

@@ -5,6 +5,7 @@ import (
"encoding/json" "encoding/json"
"errors" "errors"
"fmt" "fmt"
"reflect"
"slices" "slices"
"strings" "strings"
"testing" "testing"
@@ -60,7 +61,10 @@ func (m *fakeMTPModel) Forward(b *batch.Batch, caches []cache.Cache) *mlx.Array
mlx.Eval(b.InputIDs) mlx.Eval(b.InputIDs)
ids := b.InputIDs.Ints() ids := b.InputIDs.Ints()
m.forwards = append(m.forwards, forwardCall{offset: b.SeqOffsets[0], n: int32(len(ids))}) m.forwards = append(m.forwards, forwardCall{offset: b.SeqOffsets[0], n: int32(len(ids))})
for _, c := range caches { for i, c := range caches {
if i >= m.NumLayers() {
break
}
if rc, ok := c.(*fakeRewindableCache); ok { if rc, ok := c.(*fakeRewindableCache); ok {
seg := make([]int32, len(ids)) seg := make([]int32, len(ids))
for i, id := range ids { for i, id := range ids {
@@ -85,25 +89,10 @@ func (m *fakeMTPModel) LoadWeights(map[string]*mlx.Array) error {
return nil return nil
} }
// TokenEmbeddings returns a width-1 embedding holding the token id as a float, var _ base.Model = (*fakeMTPModel)(nil)
// so a draft can recover which token it is extending from inputs[...,0].
func (m *fakeMTPModel) TokenEmbeddings(inputIDs *mlx.Array) *mlx.Array {
mlx.Eval(inputIDs)
ids := inputIDs.Ints()
data := make([]float32, len(ids))
for i, id := range ids {
data[i] = float32(id)
}
return mlx.FromValues(data, inputIDs.Dim(0), inputIDs.Dim(1), 1)
}
var ( // fakeMTPDraft is a cacheless draft that extends b.InputIDs through predict;
_ base.Model = (*fakeMTPModel)(nil) // a map (not a step counter) keeps drafting consistent regardless of batching.
_ base.MTPEmbeddingModel = (*fakeMTPModel)(nil)
)
// fakeMTPDraft extends the token in inputEmbeds through predict; a map (not a
// step counter) keeps drafting consistent regardless of batching.
type fakeMTPDraft struct { type fakeMTPDraft struct {
predict map[int32]int32 predict map[int32]int32
// calls records each Draft call so tests can assert the position convention. // calls records each Draft call so tests can assert the position convention.
@@ -117,17 +106,82 @@ type draftCall struct {
func (d *fakeMTPDraft) LoadWeights(map[string]*mlx.Array) error { return nil } func (d *fakeMTPDraft) LoadWeights(map[string]*mlx.Array) error { return nil }
func (d *fakeMTPDraft) Draft(inputEmbeds *mlx.Array, position int32, caches []cache.Cache) (logits, hidden *mlx.Array) { func (d *fakeMTPDraft) DraftCaches([]cache.Cache) []cache.Cache { return nil }
mlx.Eval(inputEmbeds)
prev := int32(inputEmbeds.Floats()[0]) func (d *fakeMTPDraft) Draft(b *batch.Batch, caches []cache.Cache) (hidden, projected *mlx.Array) {
d.calls = append(d.calls, draftCall{position: position, from: prev}) mlx.Eval(b.InputIDs)
prev := int32(b.InputIDs.Ints()[0])
d.calls = append(d.calls, draftCall{position: b.SeqOffsets[0], from: prev})
return oneHotLogits([]int32{d.predict[prev]}), mlx.Zeros(mlx.DTypeFloat32, 1, 1, mtpTestVocab) return oneHotLogits([]int32{d.predict[prev]}), mlx.Zeros(mlx.DTypeFloat32, 1, 1, mtpTestVocab)
} }
var ( // Unembed is the identity: the fake's hidden already is its one-hot logits.
_ base.DraftModel = (*fakeMTPDraft)(nil) func (d *fakeMTPDraft) Unembed(x *mlx.Array) *mlx.Array { return x }
_ base.MTPDraftModel = (*fakeMTPDraft)(nil)
) var _ base.DraftModel = (*fakeMTPDraft)(nil)
// fakeKVDraft is a draft head with its own KV cache: it claims the trailing
// cache slot, writes its input ids there on every Draft call (advancing the
// offset like a real KV write), and records each call's offset, ids, and
// the identity of every fused hidden row. A target hidden row is one-hot
// (its hot index identifies which position it came from); the head's own
// projected hidden is all-zero and records as -1.
type fakeKVDraft struct {
predict map[int32]int32
extends []extendCall
}
// extendCall is one recorded Draft call: the absolute slot of the first
// entry written, the look-ahead token ids, and the hot index of each fused
// hidden row (-1 for the head's own projected hidden).
type extendCall struct {
offset int32
ids []int32
hiddens []int32
}
func (d *fakeKVDraft) LoadWeights(map[string]*mlx.Array) error { return nil }
func (d *fakeKVDraft) DraftCaches(caches []cache.Cache) []cache.Cache {
return caches[len(caches)-1:]
}
func (d *fakeKVDraft) Draft(b *batch.Batch, caches []cache.Cache) (hidden, projected *mlx.Array) {
mlx.Eval(b.InputIDs, b.Hidden)
rawIDs := b.InputIDs.Ints()
ids := make([]int32, len(rawIDs))
for i, id := range rawIDs {
ids[i] = int32(id)
}
hot := make([]int32, b.Hidden.Dim(1))
flat := b.Hidden.Floats()
for r := range hot {
hot[r] = -1
for v := range mtpTestVocab {
if flat[r*mtpTestVocab+v] != 0 {
hot[r] = int32(v)
break
}
}
}
d.extends = append(d.extends, extendCall{offset: b.SeqOffsets[0], ids: ids, hiddens: hot})
if rc, ok := d.DraftCaches(caches)[0].(*fakeRewindableCache); ok {
rc.feed(ids)
}
preds := make([]int32, len(ids))
for i, id := range ids {
preds[i] = d.predict[id]
}
return oneHotLogits(preds), mlx.Zeros(mlx.DTypeFloat32, 1, 1, mtpTestVocab)
}
// Unembed is the identity: the fake's hidden already is its one-hot logits.
func (d *fakeKVDraft) Unembed(x *mlx.Array) *mlx.Array { return x }
var _ base.DraftModel = (*fakeKVDraft)(nil)
// newTestTokenizer builds a byte-level BPE tokenizer over single-character // newTestTokenizer builds a byte-level BPE tokenizer over single-character
// tokens "0".."7" with the given EOS ids, so Decode(id) yields that digit and // tokens "0".."7" with the given EOS ids, so Decode(id) yields that digit and
@@ -215,7 +269,7 @@ func TestAcceptMTPDraftsGreedyAcceptAll(t *testing.T) {
spec := testSpeculationSession(r, caches) spec := testSpeculationSession(r, caches)
current := sampler.Result{Token: mlx.FromValues([]int32{1}, 1)} current := sampler.Result{Token: mlx.FromValues([]int32{1}, 1)}
results, accepted, err := spec.accept(&position, current, oneHotLogits([]int32{2}), baseLogits, candidates) results, accepted, err := spec.accept(&position, current, oneHotLogits([]int32{1}), baseLogits, candidates)
if err != nil { if err != nil {
t.Fatalf("accept: %v", err) t.Fatalf("accept: %v", err)
} }
@@ -250,7 +304,7 @@ func TestAcceptMTPDraftsGreedyMismatch(t *testing.T) {
spec := testSpeculationSession(r, caches) spec := testSpeculationSession(r, caches)
current := sampler.Result{Token: mlx.FromValues([]int32{1}, 1)} current := sampler.Result{Token: mlx.FromValues([]int32{1}, 1)}
results, accepted, err := spec.accept(&position, current, oneHotLogits([]int32{2}), baseLogits, candidates) results, accepted, err := spec.accept(&position, current, oneHotLogits([]int32{1}), baseLogits, candidates)
if err != nil { if err != nil {
t.Fatalf("accept: %v", err) t.Fatalf("accept: %v", err)
} }
@@ -273,8 +327,8 @@ func TestAcceptMTPDraftsGreedyMismatch(t *testing.T) {
func TestAcceptMTPDraftsGreedyEOS(t *testing.T) { func TestAcceptMTPDraftsGreedyEOS(t *testing.T) {
skipIfNoMLX(t) skipIfNoMLX(t)
// The second accepted draft token is EOS: it is recorded but stops // The second accepted draft token is EOS: it is recorded but stops
// generation and no bonus token is produced. Both accepted tokens are // generation and no bonus token is produced. The EOS's own KV is rolled
// committed, so the caches hold the EOS's KV. // back so the caches rest one token behind the recorded outputs.
const eos int32 = 6 const eos int32 = 6
predict := map[int32]int32{1: 2, 2: eos, eos: 0} predict := map[int32]int32{1: 2, 2: eos, eos: 0}
r := mtpTestRunner(t, predict, []int32{eos}, sampler.Options{}) r := mtpTestRunner(t, predict, []int32{eos}, sampler.Options{})
@@ -287,7 +341,7 @@ func TestAcceptMTPDraftsGreedyEOS(t *testing.T) {
spec := testSpeculationSession(r, caches) spec := testSpeculationSession(r, caches)
current := sampler.Result{Token: mlx.FromValues([]int32{1}, 1)} current := sampler.Result{Token: mlx.FromValues([]int32{1}, 1)}
results, accepted, err := spec.accept(&position, current, oneHotLogits([]int32{2}), baseLogits, candidates) results, accepted, err := spec.accept(&position, current, oneHotLogits([]int32{1}), baseLogits, candidates)
if err != nil { if err != nil {
t.Fatalf("accept: %v", err) t.Fatalf("accept: %v", err)
} }
@@ -302,11 +356,11 @@ func TestAcceptMTPDraftsGreedyEOS(t *testing.T) {
if len(results) != accepted { if len(results) != accepted {
t.Fatalf("results has %d tokens, want exactly the %d accepted with no bonus after EOS", len(results), accepted) t.Fatalf("results has %d tokens, want exactly the %d accepted with no bonus after EOS", len(results), accepted)
} }
if position != 2 { if position != 1 {
t.Fatalf("position = %d, want 2 (both accepted tokens committed)", position) t.Fatalf("position = %d, want 1 (kept draft committed, EOS dropped)", position)
} }
if got := caches[0].Offset(); got != 2 { if got := caches[0].Offset(); got != 1 {
t.Fatalf("cache offset = %d, want 2 (both accepted tokens committed)", got) t.Fatalf("cache offset = %d, want 1 (one behind the recorded outputs, EOS dropped)", got)
} }
} }
@@ -392,7 +446,11 @@ func TestRunMTPDecodeSampled(t *testing.T) {
CompletionRequest: CompletionRequest{Options: api.Options{NumPredict: 20}}, CompletionRequest: CompletionRequest{Options: api.Options{NumPredict: 20}},
SamplerOpts: sampler.Options{Temperature: 1, Seed: 42, UseSeed: true}, SamplerOpts: sampler.Options{Temperature: 1, Seed: 42, UseSeed: true},
} }
d := testDecoder(r, req, caches, []int32{1}, position) spec := r.spec.open(req, caches)
if spec == nil || !spec.enabled {
t.Fatalf("newSpeculationSession rejected a sampled request")
}
d := spec.decoder([]int32{1}, position)
if err := r.decode(context.Background(), req, session, d, 0); err != nil { if err := r.decode(context.Background(), req, session, d, 0); err != nil {
t.Fatalf("decode: %v", err) t.Fatalf("decode: %v", err)
} }
@@ -518,7 +576,299 @@ func testDecoder(r *Runner, req Request, caches []cache.Cache, seed []int32, pos
if spec := r.spec.open(req, caches); spec != nil { if spec := r.spec.open(req, caches); spec != nil {
return spec.decoder(seed, position) return spec.decoder(seed, position)
} }
return r.pipelinedDecoder(caches, seed, position) return r.pipelinedDecoder(nil, caches, seed, position)
}
func TestDecodeKVDraft(t *testing.T) {
skipIfNoMLX(t)
// A draft with its own KV cache mirroring the target chain
// 1->2->3->4->5->6->EOS.
// The decode seed catch-up writes the draft pair for the prompt token,
// the first proposal comes from the catch-up's held logits (no draft
// call), speculative entries are written at advancing slots, and the
// post-accept rebuild rewrites the committed range from target hiddens.
const eos int32 = 7
predict := map[int32]int32{1: 2, 2: 3, 3: 4, 4: 5, 5: 6, 6: eos, eos: 0}
r := mtpTestRunner(t, predict, []int32{eos}, sampler.Options{})
draft := &fakeKVDraft{predict: predict}
r.spec = newSpeculation(r, draft)
caches, _ := newMTPTestCaches(2) // caches[0] target, caches[1] draft KV
session, ch := newMTPTestSession(caches)
position := 0
req := Request{
Responses: ch,
Tokens: []int32{1},
CompletionRequest: CompletionRequest{Options: api.Options{NumPredict: 20}},
SamplerOpts: sampler.Options{},
}
spec := r.spec.open(req, caches)
if spec == nil || !spec.enabled || len(spec.spec.targets) != 1 {
t.Fatalf("speculation engine not built around the draft caches")
}
defer spec.close()
d := spec.decoder([]int32{1}, position)
if err := r.decode(context.Background(), req, session, d, 0); err != nil {
t.Fatalf("decode: %v", err)
}
d.close()
content, final := collectResponses(ch)
if content != "23456" {
t.Fatalf("content = %q, want %q", content, "23456")
}
if final.EvalCount != 5 || final.DoneReason != 0 {
t.Fatalf("EvalCount = %d, DoneReason = %d, want 5 and 0", final.EvalCount, final.DoneReason)
}
if got := []int32{2, 3, 4, 5, 6, eos}; !slices.Equal(session.outputs, got) {
t.Fatalf("session outputs = %v, want %v", session.outputs, got)
}
// All caches rest together at the trie frontier: the prompt token plus
// five generated tokens; the EOS's KV is never committed.
if got := caches[0].Offset(); got != 6 {
t.Fatalf("target offset = %d, want 6", got)
}
if got := caches[1].Offset(); got != 6 {
t.Fatalf("draft cache offset = %d, want 6 (lockstep with target)", got)
}
// Every committed draft slot S fuses look-ahead token x_{S+1} with the
// target hidden at S (whose hot index equals the look-ahead id on this
// mirrored chain); speculative steps fuse the head's own projections
// (-1), seeded by the catch-up flush's held projection. The first
// proposal of the round consumes the catch-up's held logits, so the
// four-token draft makes only three head calls before the rebuild.
wantExtends := []extendCall{
{offset: 0, ids: []int32{2}, hiddens: []int32{2}}, // frontier pair flushed at the first proposal
{offset: 1, ids: []int32{3}, hiddens: []int32{-1}}, // speculative step 2 (held projection)
{offset: 2, ids: []int32{4}, hiddens: []int32{-1}}, // speculative step 3 (projection)
{offset: 3, ids: []int32{5}, hiddens: []int32{-1}}, // speculative step 4 (projection)
{offset: 1, ids: []int32{3, 4, 5, 6, eos}, hiddens: []int32{3, 4, 5, 6, eos}}, // committed pairs from the validated run + finish
}
if !reflect.DeepEqual(draft.extends, wantExtends) {
t.Fatalf("draft extends = %+v, want %+v", draft.extends, wantExtends)
}
// The committed draft KV holds the look-ahead tokens, one per slot.
if got, want := caches[1].(*fakeRewindableCache).tokens, []int32{2, 3, 4, 5, 6, eos}; !slices.Equal(got, want) {
t.Fatalf("draft cache = %v, want %v", got, want)
}
}
func TestDecodeKVDraftRejectionRebuildsFromTarget(t *testing.T) {
skipIfNoMLX(t)
// The draft mispredicts mid-chain: it proposes 6 where the target's own
// next token is 4, so the round is accepted only up to the rejection and the
// loop re-proposes from the target's correction. The speculative draft KV
// written for the rejected proposals is rolled back and rewritten from the
// validated run's target hiddens, so the committed draft cache holds the
// target chain with no trace of the rejected tokens.
const eos int32 = 7
target := map[int32]int32{1: 2, 2: 3, 3: 4, 4: 5, 5: eos, eos: 0}
r := mtpTestRunner(t, target, []int32{eos}, sampler.Options{})
// The draft mirrors the target except at 3, where it proposes 6 (absent from
// the target chain); once the target corrects 3->4 the next proposal
// re-aligns on the shared chain.
draft := &fakeKVDraft{predict: map[int32]int32{1: 2, 2: 3, 3: 6, 6: 0, 4: 5, 5: eos, eos: 0}}
r.spec = newSpeculation(r, draft)
caches, _ := newMTPTestCaches(2) // caches[0] target, caches[1] draft KV
session, ch := newMTPTestSession(caches)
position := 0
req := Request{
Responses: ch,
Tokens: []int32{1},
CompletionRequest: CompletionRequest{Options: api.Options{NumPredict: 20}},
SamplerOpts: sampler.Options{},
}
spec := r.spec.open(req, caches)
spec.limit = 4
defer spec.close()
d := spec.decoder([]int32{1}, position)
if err := r.decode(context.Background(), req, session, d, 0); err != nil {
t.Fatalf("decode: %v", err)
}
d.close()
content, final := collectResponses(ch)
if content != "2345" {
t.Fatalf("content = %q, want %q", content, "2345")
}
if final.DoneReason != 0 {
t.Fatalf("DoneReason = %d, want 0 (EOS)", final.DoneReason)
}
if got := []int32{2, 3, 4, 5, eos}; !slices.Equal(session.outputs, got) {
t.Fatalf("session outputs = %v, want %v", session.outputs, got)
}
// The divergent token was drafted into the speculative KV...
drafted := false
for _, e := range draft.extends {
if slices.Contains(e.ids, 6) {
drafted = true
}
}
if !drafted {
t.Fatalf("draft never proposed the divergent token 6; extends = %+v", draft.extends)
}
// ...but the rejection rolled it back: both caches rest in lockstep at the
// target chain, and the committed draft KV holds only the validated
// look-ahead tokens.
if got := caches[0].Offset(); got != 5 {
t.Fatalf("target offset = %d, want 5", got)
}
if got := caches[1].Offset(); got != 5 {
t.Fatalf("draft cache offset = %d, want 5 (lockstep with target)", got)
}
if got, want := caches[1].(*fakeRewindableCache).tokens, []int32{2, 3, 4, 5, eos}; !slices.Equal(got, want) {
t.Fatalf("draft cache = %v, want %v (rejected proposals rolled back)", got, want)
}
}
func TestDecodeMaintainsDraftCacheWithoutDrafting(t *testing.T) {
skipIfNoMLX(t)
// A request that cannot speculate (logprobs) on a model whose draft has
// its own KV cache still maintains it: the pipelined decoder
// reports each forwarded token, the pairs wait in the pending list, and
// one batched extend at close — its final pair completed by the decoder's
// discarded in-flight sample — leaves the draft level with the target.
const eos int32 = 7
predict := map[int32]int32{1: 2, 2: 3, 3: 4, 4: 5, 5: 6, 6: eos, eos: 0}
opts := sampler.Options{Logprobs: true}
r := mtpTestRunner(t, predict, []int32{eos}, opts)
draft := &fakeKVDraft{predict: predict}
r.spec = newSpeculation(r, draft)
caches, _ := newMTPTestCaches(2)
session, ch := newMTPTestSession(caches)
position := 0
req := Request{
Responses: ch,
Tokens: []int32{1},
CompletionRequest: CompletionRequest{Options: api.Options{NumPredict: 20}},
SamplerOpts: opts,
}
spec := r.spec.open(req, caches)
if spec == nil || spec.enabled {
t.Fatalf("want a maintain-only speculationSession, got %+v", spec)
}
d := spec.decoder([]int32{1}, position)
if err := r.decode(context.Background(), req, session, d, 0); err != nil {
t.Fatalf("decode: %v", err)
}
d.close()
spec.close() // the pipeline defers this before session close
content, _ := collectResponses(ch)
if content != "23456" {
t.Fatalf("content = %q, want %q", content, "23456")
}
// The pipelined decoder forwards every emitted token, including the
// ending EOS, so both caches rest at the full recorded path.
if got := caches[0].Offset(); got != 7 {
t.Fatalf("target offset = %d, want 7", got)
}
if got := caches[1].Offset(); got != 7 {
t.Fatalf("draft cache offset = %d, want 7 (lockstep with target)", got)
}
// A session that never drafts defers every committed pair: the whole
// generation arrives in one batched flush at finish, with the EOS slot's
// pair completed by the in-flight sample (predict[eos] = 0) that decoding
// discarded.
wantExtends := []extendCall{
{offset: 0, ids: []int32{2, 3, 4, 5, 6, eos, 0}, hiddens: []int32{2, 3, 4, 5, 6, eos, 0}},
}
if !reflect.DeepEqual(draft.extends, wantExtends) {
t.Fatalf("draft extends = %+v, want %+v", draft.extends, wantExtends)
}
if got, want := caches[1].(*fakeRewindableCache).tokens, []int32{2, 3, 4, 5, 6, eos, 0}; !slices.Equal(got, want) {
t.Fatalf("draft cache = %v, want %v", got, want)
}
}
func TestFlushLevelsDraftCacheWithPrefill(t *testing.T) {
skipIfNoMLX(t)
// Prefill attaches its scheduled snapshots only at offsets every cache
// has crossed, so the pipeline flushes the drafter's buffered pairs
// first: after a prefill-sized committed report and a flush, the
// draft cache covers every completed pair, one slot behind the target
// (the frontier pair still awaits its look-ahead token).
const eos int32 = 7
predict := map[int32]int32{2: 3, 3: 4, 4: 5}
r := mtpTestRunner(t, predict, []int32{eos}, sampler.Options{})
draft := &fakeKVDraft{predict: predict}
r.spec = newSpeculation(r, draft)
caches, _ := newMTPTestCaches(2)
req := Request{
CompletionRequest: CompletionRequest{Options: api.Options{NumPredict: 20}},
SamplerOpts: sampler.Options{},
}
spec := r.spec.open(req, caches)
defer spec.close()
// The prompt's only chunk: tokens 1..4 at slots 0..3 with their hiddens.
spec.committed(mlx.FromValues([]int32{1, 2, 3, 4}, 1, 4), oneHotLogits([]int32{1, 2, 3, 4}), 0)
spec.flush()
if got := caches[1].Offset(); got != 3 {
t.Fatalf("draft cache offset after flush = %d, want 3 (all completed pairs)", got)
}
wantExtends := []extendCall{
{offset: 0, ids: []int32{2, 3, 4}, hiddens: []int32{1, 2, 3}},
}
if !reflect.DeepEqual(draft.extends, wantExtends) {
t.Fatalf("draft extends = %+v, want %+v", draft.extends, wantExtends)
}
}
func TestCommittedRunBatchesPastFlushCap(t *testing.T) {
skipIfNoMLX(t)
// A committed run longer than the pending-flush cap still writes the draft
// caches in a single head forward: the run's completed pairs coalesce into
// one batched extend at the run's start, rather than splitting at the cap.
const n = mtpPendingFlushTokens + 8
predict := map[int32]int32{}
tokens := make([]int32, n)
for i := range tokens {
tokens[i] = int32(i%7 + 1)
predict[tokens[i]] = int32((i+1)%7 + 1)
}
r := mtpTestRunner(t, predict, []int32{0}, sampler.Options{})
draft := &fakeKVDraft{predict: predict}
r.spec = newSpeculation(r, draft)
caches, _ := newMTPTestCaches(2)
req := Request{
CompletionRequest: CompletionRequest{Options: api.Options{NumPredict: 20}},
SamplerOpts: sampler.Options{},
}
spec := r.spec.open(req, caches)
defer spec.close()
// One prefill-sized chunk: n tokens at slots 0..n-1 with their hiddens.
spec.committed(mlx.FromValues(tokens, 1, n), oneHotLogits(tokens), 0)
spec.flush()
if got := len(draft.extends); got != 1 {
t.Fatalf("draft extends = %d calls, want 1 batched extend", got)
}
if got, want := len(draft.extends[0].ids), n-1; got != want {
t.Fatalf("batched extend ids = %d, want %d (every completed pair)", got, want)
}
if got := draft.extends[0].offset; got != 0 {
t.Fatalf("batched extend offset = %d, want 0", got)
}
if got := caches[1].Offset(); got != n-1 {
t.Fatalf("draft cache offset = %d, want %d (all completed pairs)", got, n-1)
}
} }
// newMTPTestCaches returns n rewindable fake caches sharing one snapshot // newMTPTestCaches returns n rewindable fake caches sharing one snapshot
@@ -539,19 +889,16 @@ func newMTPTestSession(caches []cache.Cache) (*cacheSession, chan CompletionResp
return &cacheSession{caches: caches}, ch return &cacheSession{caches: caches}, ch
} }
// testSpeculationSession wires a speculation engine around the runner's drafter and // testSpeculationSession wires a speculation engine around a drafter and caches for
// caches for tests that drive accept directly. A runner without a draft // tests that drive accept directly. A runner without a draft model gets a
// model gets a no-op drafter, since the engine requires one. // no-op drafter, since the engine requires one.
func testSpeculationSession(r *Runner, caches []cache.Cache) *speculationSession { func testSpeculationSession(r *Runner, caches []cache.Cache) *speculationSession {
s := r.spec if r.spec != nil {
if s == nil { r.spec.bind(caches)
s = &speculation{r: r} return &speculationSession{spec: r.spec, drafter: newMTPDrafter(r.spec)}
} }
var d drafter = nopDrafter{} s := &speculation{r: r, caches: caches, targets: caches}
if md := newMTPDrafter(s, caches); md != nil { return &speculationSession{spec: s, drafter: nopDrafter{}}
d = md
}
return &speculationSession{spec: s, drafter: d, caches: caches}
} }
// nopDrafter satisfies drafter for engine tests that supply candidates // nopDrafter satisfies drafter for engine tests that supply candidates
@@ -560,6 +907,8 @@ type nopDrafter struct{}
func (nopDrafter) propose(*mlx.Array, int) *draftCandidates { return nil } func (nopDrafter) propose(*mlx.Array, int) *draftCandidates { return nil }
func (nopDrafter) committed(_, _ *mlx.Array, _ int) {} func (nopDrafter) committed(_, _ *mlx.Array, _ int) {}
func (nopDrafter) finish(*mlx.Array) {}
func (nopDrafter) flush() {}
func (nopDrafter) close() {} func (nopDrafter) close() {}
// scriptedCandidates builds draft candidates by running the real drafter // scriptedCandidates builds draft candidates by running the real drafter
@@ -573,7 +922,8 @@ func scriptedCandidates(r *Runner, tokens []int32) *draftCandidates {
chain[prev] = tok chain[prev] = tok
prev = tok prev = tok
} }
d := &mtpDrafter{spec: &speculation{r: r}, draft: &fakeMTPDraft{predict: chain}, target: r.Model.(base.MTPEmbeddingModel)} s := &speculation{r: r, draft: &fakeMTPDraft{predict: chain}}
d := &mtpDrafter{spec: s}
d.committed(mlx.FromValues([]int32{0}, 1, 1), mlx.Zeros(mlx.DTypeFloat32, 1, 1, mtpTestVocab), 0) d.committed(mlx.FromValues([]int32{0}, 1, 1), mlx.Zeros(mlx.DTypeFloat32, 1, 1, mtpTestVocab), 0)
return d.propose(mlx.FromValues([]int32{0}, 1), len(tokens)) return d.propose(mlx.FromValues([]int32{0}, 1), len(tokens))
} }

View File

@@ -75,7 +75,12 @@ func (r *Runner) TextGenerationPipeline(ctx context.Context, request Request) er
defer session.close() defer session.close()
caches := session.caches caches := session.caches
seed, position, promptEval, err := r.prefill(ctx, session) // Built before prefill so a drafter with draft caches follows the prompt
// through prefill alongside the target.
spec := r.spec.open(request, caches)
defer spec.close()
seed, position, promptEval, err := r.prefill(ctx, session, spec)
if err != nil { if err != nil {
return err return err
} }
@@ -83,14 +88,11 @@ func (r *Runner) TextGenerationPipeline(ctx context.Context, request Request) er
// Register the sampler after prefill completes. // Register the sampler after prefill completes.
r.Sampler.Add(pipelineSlot, request.SamplerOpts, inputs) r.Sampler.Add(pipelineSlot, request.SamplerOpts, inputs)
spec := r.spec.open(request, caches)
defer spec.close()
var d decoder var d decoder
if spec != nil { if spec != nil {
d = spec.decoder(seed, position) d = spec.decoder(seed, position)
} else { } else {
d = r.pipelinedDecoder(caches, seed, position) d = r.pipelinedDecoder(nil, caches, seed, position)
} }
defer d.close() defer d.close()
return r.decode(ctx, request, session, d, promptEval) return r.decode(ctx, request, session, d, promptEval)
@@ -100,7 +102,7 @@ func (r *Runner) TextGenerationPipeline(ctx context.Context, request Request) er
// one token for decode to seed from, and schedules the prompt's periodic // one token for decode to seed from, and schedules the prompt's periodic
// snapshots. It returns the seed tokens, the resume position, and the // snapshots. It returns the seed tokens, the resume position, and the
// prompt-evaluation duration. // prompt-evaluation duration.
func (r *Runner) prefill(ctx context.Context, session *cacheSession) ([]int32, int, time.Duration, error) { func (r *Runner) prefill(ctx context.Context, session *cacheSession, spec *speculationSession) ([]int32, int, time.Duration, error) {
start := time.Now() start := time.Now()
inputs := session.inputs inputs := session.inputs
tokens := session.remaining tokens := session.remaining
@@ -143,11 +145,13 @@ func (r *Runner) prefill(ctx context.Context, session *cacheSession) ([]int32, i
n := min(prefillChunk, total-processed-1) n := min(prefillChunk, total-processed-1)
r.Model.Forward(&batch.Batch{ chunkIDs := mlx.FromValues(tokens[processed:processed+n], 1, n)
InputIDs: mlx.FromValues(tokens[processed:processed+n], 1, n), hidden := r.Model.Forward(&batch.Batch{
InputIDs: chunkIDs,
SeqOffsets: []int32{int32(position)}, SeqOffsets: []int32{int32(position)},
SeqQueryLens: []int32{int32(n)}, SeqQueryLens: []int32{int32(n)},
}, caches) }, caches)
spec.committed(chunkIDs, hidden, position)
mlx.Sweep() mlx.Sweep()
materializeCaches() materializeCaches()
processed += n processed += n
@@ -158,7 +162,10 @@ func (r *Runner) prefill(ctx context.Context, session *cacheSession) ([]int32, i
mlx.ClearCache() mlx.ClearCache()
} }
// Attach the snapshots captured during prefill to the trie. // Flush before attaching: snapshots attach only at offsets every cache
// has crossed, and a drafter with draft caches keeps buffered pairs that
// would otherwise hold those caches short of the scheduled offsets.
spec.flush()
session.attachPrefillSnapshots() session.attachPrefillSnapshots()
return tokens[processed:], position, time.Since(start), nil return tokens[processed:], position, time.Since(start), nil
@@ -263,15 +270,18 @@ func (r *Runner) decode(ctx context.Context, request Request, session *cacheSess
// the next token's chain is dispatched before the returned one is // the next token's chain is dispatched before the returned one is
// synchronized, so the device runs ahead of host emission. // synchronized, so the device runs ahead of host emission.
type pipelinedDecoder struct { type pipelinedDecoder struct {
r *Runner r *Runner
// spec, when non-nil, receives every forwarded token and settles its
// drafter at close, keeping a non-drafting session's draft KV level.
spec *speculationSession
caches []cache.Cache caches []cache.Cache
position int position int
sample sampler.Result // in flight: sampled, not yet forwarded sample sampler.Result // in flight: sampled, not yet forwarded
emitted sampler.Result // last call's result, pinned until the next call emitted sampler.Result // last call's result, pinned until the next call
} }
func (r *Runner) pipelinedDecoder(caches []cache.Cache, seed []int32, position int) *pipelinedDecoder { func (r *Runner) pipelinedDecoder(spec *speculationSession, caches []cache.Cache, seed []int32, position int) *pipelinedDecoder {
t := &pipelinedDecoder{r: r, caches: caches, position: position} t := &pipelinedDecoder{r: r, spec: spec, caches: caches, position: position}
t.sample = t.dispatch(mlx.FromValues(seed, 1, len(seed))) t.sample = t.dispatch(mlx.FromValues(seed, 1, len(seed)))
return t return t
} }
@@ -285,8 +295,10 @@ func (t *pipelinedDecoder) dispatch(token *mlx.Array) sampler.Result {
SeqOffsets: []int32{int32(t.position)}, SeqOffsets: []int32{int32(t.position)},
SeqQueryLens: []int32{int32(token.Dim(1))}, SeqQueryLens: []int32{int32(token.Dim(1))},
}, t.caches) }, t.caches)
t.spec.committed(token, hidden, t.position)
t.position += token.Dim(1) t.position += token.Dim(1)
next := r.Sampler.Sample([]int{pipelineSlot}, lastLogits(r.Model.Unembed(hidden))) logits := r.Model.Unembed(hidden)
next := r.Sampler.Sample([]int{pipelineSlot}, logits.Slice(mlx.Slice(), mlx.Slice(logits.Dim(1)-1), mlx.Slice()).Squeeze(1))
mlx.Pin(next.Arrays()...) mlx.Pin(next.Arrays()...)
mlx.Sweep() mlx.Sweep()
mlx.AsyncEval(next.Arrays()...) mlx.AsyncEval(next.Arrays()...)
@@ -300,6 +312,9 @@ func (t *pipelinedDecoder) next(int) ([]sampler.Result, error) {
} }
func (t *pipelinedDecoder) close() { func (t *pipelinedDecoder) close() {
// The in-flight sample's forward was never dispatched; its report lets
// the drafter settle level with the caches' resting offset.
t.spec.finish(t.sample.Token)
mlx.Unpin(t.emitted.Arrays()...) mlx.Unpin(t.emitted.Arrays()...)
mlx.Unpin(t.sample.Arrays()...) mlx.Unpin(t.sample.Arrays()...)
} }

View File

@@ -81,6 +81,9 @@ func (r *Runner) Load(modelName string) error {
return err return err
} }
draftModel = draft draftModel = draft
} else if sd, ok := m.(base.SelfDraft); ok {
// Inline draft head: already loaded with the target; nil if none shipped.
draftModel = sd.SelfDraft()
} }
collected := mlx.Collect(m) collected := mlx.Collect(m)

View File

@@ -4,6 +4,7 @@ import (
"fmt" "fmt"
"log/slog" "log/slog"
"os" "os"
"slices"
"strconv" "strconv"
"strings" "strings"
@@ -23,10 +24,18 @@ type drafter interface {
// committed reports a run of tokens committed to the target caches: // committed reports a run of tokens committed to the target caches:
// tokens[i] sits at slot position+i and hiddens row i is the target // tokens[i] sits at slot position+i and hiddens row i is the target
// hidden state at that slot. Runs arrive in slot order — the decode // hidden state at that slot. Runs arrive in slot order — prefill
// seed, then each round's validated tokens. // chunks, the decode seed, then each round's validated tokens.
committed(tokens, hiddens *mlx.Array, position int) committed(tokens, hiddens *mlx.Array, position int)
// finish reports generation ended with current sampled but never
// committed, so the drafter can settle state tracking the target caches.
finish(current *mlx.Array)
// flush writes any buffered committed reports through to the draft
// caches. A drafter without draft caches has nothing to write.
flush()
close() close()
} }
@@ -100,6 +109,15 @@ func positiveEnvInt(key string) int {
type speculation struct { type speculation struct {
r *Runner r *Runner
draft base.DraftModel draft base.DraftModel
// caches is the whole persistent slice, passed to every forward; draftKV
// are the draft head's own caches and targets are the rest — the caches the
// target forward writes, which speculation snapshots and rollback cover.
// Bound the first time the caches exist (the Runner reuses one cache slice
// for its life) and stable thereafter.
caches []cache.Cache
draftKV []cache.Cache
targets []cache.Cache
} }
// newSpeculation builds the speculative-decoding subsystem for a loaded model, // newSpeculation builds the speculative-decoding subsystem for a loaded model,
@@ -111,41 +129,68 @@ func newSpeculation(r *Runner, draft base.DraftModel) *speculation {
return &speculation{r: r, draft: draft} return &speculation{r: r, draft: draft}
} }
// bind computes the draft/target cache partition the first time the persistent
// caches exist; later requests reuse the same slice, so it runs once.
func (s *speculation) bind(caches []cache.Cache) {
if s.caches != nil {
if !slices.Equal(s.caches, caches) {
panic("speculation: cache slice changed between requests")
}
return
}
draftKV := s.draft.DraftCaches(caches)
// Partition caches into target slots (everything not in draftKV) in one
// pass. The count check rejects a draft slot that isn't a member of caches.
targets := make([]cache.Cache, 0, len(caches))
for _, c := range caches {
if !slices.Contains(draftKV, c) {
targets = append(targets, c)
}
}
if len(caches)-len(targets) != len(draftKV) {
panic("speculation: DraftCaches must select slots of the cache slice")
}
s.caches = caches
s.draftKV = draftKV
s.targets = targets
}
// speculationSession is the speculation cursor for one request over a speculation. A // speculationSession is the speculation cursor for one request over a speculation. A
// nil speculationSession is a plain decode. // nil speculationSession is a plain decode.
type speculationSession struct { type speculationSession struct {
spec *speculation spec *speculation
drafter drafter drafter drafter
caches []cache.Cache // enabled selects speculative rounds; a maintain-only engine decodes
// plainly while streaming committed runs to keep draft KV prefix-cached.
enabled bool
opts specOptions opts specOptions
limit int // current draft length limit int // current draft length
stats specStats stats specStats
} }
// open returns the speculation cursor for this request, or nil when the model // open returns the speculation cursor for this request or nil when the model ships
// carries no draft head or the request cannot speculate. Logprobs are not yet // no draft head (a nil receiver), which decodes plainly.
// supported. A nil receiver (no draft head) decodes plainly.
func (s *speculation) open(request Request, caches []cache.Cache) *speculationSession { func (s *speculation) open(request Request, caches []cache.Cache) *speculationSession {
if s == nil { if s == nil {
return nil return nil
} }
d := newMTPDrafter(s, caches) s.bind(caches)
d := newMTPDrafter(s)
if d == nil { if d == nil {
return nil return nil
} }
opts := request.SamplerOpts opts := request.SamplerOpts
if !s.r.mtpDefaults(opts.Temperature != 0).Enabled { enabled := s.r.mtpDefaults(opts.Temperature != 0).Enabled &&
return nil !opts.Logprobs && opts.TopLogprobs == 0
}
if opts.Logprobs || opts.TopLogprobs > 0 {
return nil
}
specOpts := s.r.loadSpecOptions(opts.Temperature != 0) specOpts := s.r.loadSpecOptions(opts.Temperature != 0)
return &speculationSession{ return &speculationSession{
spec: s, spec: s,
drafter: d, drafter: d,
caches: caches, enabled: enabled,
opts: specOpts, opts: specOpts,
limit: specOpts.initialDraftTokens, limit: specOpts.initialDraftTokens,
stats: specStats{maxDraft: specOpts.initialDraftTokens}, stats: specStats{maxDraft: specOpts.initialDraftTokens},
@@ -159,6 +204,23 @@ func (s *speculationSession) committed(tokens, hiddens *mlx.Array, position int)
s.drafter.committed(tokens, hiddens, position) s.drafter.committed(tokens, hiddens, position)
} }
// finish reports the end of generation to the drafter: current was sampled
// after the last committed slot and will never be committed.
func (s *speculationSession) finish(current *mlx.Array) {
if s == nil {
return
}
s.drafter.finish(current)
}
// flush writes the drafter's buffered committed reports to the draft caches.
func (s *speculationSession) flush() {
if s == nil {
return
}
s.drafter.flush()
}
func (s *speculationSession) close() { func (s *speculationSession) close() {
if s == nil { if s == nil {
return return
@@ -178,7 +240,13 @@ type speculativeDecoder struct {
current sampler.Result // emitted (or the seed), not yet forwarded current sampler.Result // emitted (or the seed), not yet forwarded
} }
func (s *speculationSession) decoder(seed []int32, position int) *speculativeDecoder { // decoder returns the decoder for this engine's session. A maintain-only
// engine decodes plainly via the pipelined decoder, which reports every
// forwarded token to keep draft KV level with the target.
func (s *speculationSession) decoder(seed []int32, position int) decoder {
if !s.enabled {
return s.spec.r.pipelinedDecoder(s, s.spec.caches, seed, position)
}
current := sampler.Result{Token: mlx.FromValues(seed, len(seed))} current := sampler.Result{Token: mlx.FromValues(seed, len(seed))}
mlx.Pin(current.Arrays()...) mlx.Pin(current.Arrays()...)
return &speculativeDecoder{s: s, position: position, current: current} return &speculativeDecoder{s: s, position: position, current: current}
@@ -192,7 +260,7 @@ func (st *speculativeDecoder) next(remaining int) ([]sampler.Result, error) {
InputIDs: tokenInput(st.current.Token), InputIDs: tokenInput(st.current.Token),
SeqOffsets: []int32{int32(st.position)}, SeqOffsets: []int32{int32(st.position)},
SeqQueryLens: []int32{1}, SeqQueryLens: []int32{1},
}, s.caches) }, s.spec.caches)
st.position++ st.position++
results, next, err := s.round(&st.position, st.current, hidden, remaining) results, next, err := s.round(&st.position, st.current, hidden, remaining)
@@ -212,6 +280,10 @@ func (st *speculativeDecoder) next(remaining int) ([]sampler.Result, error) {
} }
func (st *speculativeDecoder) close() { func (st *speculativeDecoder) close() {
// Generation always ends with a final token that was emitted but never
// forwarded; its report lets the drafter settle level with the caches'
// resting offset.
st.s.finish(st.current.Token)
mlx.Unpin(st.current.Arrays()...) mlx.Unpin(st.current.Arrays()...)
st.s.logStats() st.s.logStats()
} }
@@ -368,7 +440,7 @@ func (s *speculationSession) accept(position *int, current sampler.Result, hidde
r := s.spec.r r := s.spec.r
before := *position before := *position
draftCount := candidates.tokens.Dim(1) draftCount := candidates.tokens.Dim(1)
scheduleSpeculation(s.caches, before, draftCount) scheduleSpeculation(s.spec.targets, before, draftCount)
// Every exit between schedule and commit must drain the snapshot // Every exit between schedule and commit must drain the snapshot
// schedule and roll the speculative writes back out of the live caches: // schedule and roll the speculative writes back out of the live caches:
@@ -380,7 +452,7 @@ func (s *speculationSession) accept(position *int, current sampler.Result, hidde
return return
} }
committed = true committed = true
commitSpeculation(s.caches, keep, draftCount, before) commitSpeculation(s.spec.targets, keep, draftCount, before)
} }
defer commit(0) defer commit(0)
@@ -388,7 +460,7 @@ func (s *speculationSession) accept(position *int, current sampler.Result, hidde
InputIDs: candidates.tokens, InputIDs: candidates.tokens,
SeqOffsets: []int32{int32(before)}, SeqOffsets: []int32{int32(before)},
SeqQueryLens: []int32{int32(draftCount)}, SeqQueryLens: []int32{int32(draftCount)},
}, s.caches) }, s.spec.caches)
targetDist := r.Sampler.Distribution(pipelineSlot, validationLogits(r, baseLogits, hiddenSeq), candidates.tokens) targetDist := r.Sampler.Distribution(pipelineSlot, validationLogits(r, baseLogits, hiddenSeq), candidates.tokens)
draftDist := candidates.dist draftDist := candidates.dist
@@ -419,18 +491,24 @@ func (s *speculationSession) accept(position *int, current sampler.Result, hidde
break break
} }
} }
// The final token of a generation is recorded and streamed but its KV
// is never committed: the caches rest at the trie frontier.
keep := accepted
if done {
keep--
}
commit(keep)
*position = before + keep
commit(accepted) // Report the validated run — current plus the kept drafts, with the
*position = before + accepted
// Report the validated run — current plus the accepted drafts, with the
// hidden state at each token's own slot — to the drafter before // hidden state at each token's own slot — to the drafter before
// returning, so even a cancelled emission leaves the drafter's state // returning, so even a cancelled emission leaves the drafter's state
// describing exactly what the target caches hold. // describing exactly what the target caches hold. A done generation's
runIDs := append([]int32{int32(current.Token.Int())}, commitIDs...) // final token is never committed; it reaches the drafter through finish.
runIDs := append([]int32{int32(current.Token.Int())}, commitIDs[:keep]...)
runHiddens := lastHiddenRow(hidden) runHiddens := lastHiddenRow(hidden)
if accepted > 0 { if keep > 0 {
runHiddens = runHiddens.Concatenate(1, hiddenSeq.Slice(mlx.Slice(), mlx.Slice(0, accepted), mlx.Slice())) runHiddens = runHiddens.Concatenate(1, hiddenSeq.Slice(mlx.Slice(), mlx.Slice(0, keep), mlx.Slice()))
} }
s.drafter.committed(mlx.FromValues(runIDs, 1, len(runIDs)), runHiddens, before-1) s.drafter.committed(mlx.FromValues(runIDs, 1, len(runIDs)), runHiddens, before-1)

View File

@@ -12,10 +12,7 @@ import (
"github.com/ollama/ollama/x/models/nn" "github.com/ollama/ollama/x/models/nn"
) )
var ( var _ base.DraftModel = (*AssistantModel)(nil)
_ base.DraftModel = (*AssistantModel)(nil)
_ base.MTPDraftModel = (*AssistantModel)(nil)
)
type AssistantConfig struct { type AssistantConfig struct {
TextConfig TextConfig `json:"text_config"` TextConfig TextConfig `json:"text_config"`
@@ -37,6 +34,10 @@ type AssistantModel struct {
NormScaled *mlx.Array NormScaled *mlx.Array
// target supplies the scaled token embeddings the MTP head fuses with
// the target hidden state.
target *Model
*AssistantConfig *AssistantConfig
tensorPrefix string tensorPrefix string
@@ -145,6 +146,7 @@ func newAssistantModel(root *model.Root, target base.Model) (base.DraftModel, er
m := &AssistantModel{ m := &AssistantModel{
AssistantConfig: &cfg, AssistantConfig: &cfg,
tensorPrefix: tensorPrefix, tensorPrefix: tensorPrefix,
target: targetGemma,
Layers: make([]*AssistantLayer, cfg.TextConfig.NumHiddenLayers), Layers: make([]*AssistantLayer, cfg.TextConfig.NumHiddenLayers),
TensorQuant: root.AllTensorQuant(), TensorQuant: root.AllTensorQuant(),
} }
@@ -267,26 +269,42 @@ func (m *AssistantModel) precomputeScaledWeights() {
} }
} }
func (m *AssistantModel) Draft(inputsEmbeds *mlx.Array, position int32, caches []cache.Cache) (logits, hidden *mlx.Array) { // DraftCaches returns nil: the assistant keeps no KV and re-attends the
// target's caches read-only each step.
func (m *AssistantModel) DraftCaches([]cache.Cache) []cache.Cache { return nil }
// Draft is single-position: the head drafts every token as if it sat at the
// last target-seen position, so it anchors RoPE and the mask there — from the
// non-moving target full-attention cache — regardless of the advancing offset
// in b. The input fuses the target's scaled token embedding with b.Hidden.
func (m *AssistantModel) Draft(b *batch.Batch, caches []cache.Cache) (hidden, projected *mlx.Array) {
inputsEmbeds := m.target.TokenEmbeddings(b.InputIDs).Concatenate(-1, b.Hidden)
dims := inputsEmbeds.Dims() dims := inputsEmbeds.Dims()
B, L := int32(dims[0]), int32(dims[1]) B, L := int32(dims[0]), int32(dims[1])
b := &batch.Batch{
anchor := int32(0)
if len(caches) > 0 {
if full := caches[len(caches)-1]; full != nil {
anchor = int32(full.Offset() - 1)
}
}
ab := &batch.Batch{
InputIDs: mlx.Zeros(mlx.DTypeInt32, int(B), int(L)), InputIDs: mlx.Zeros(mlx.DTypeInt32, int(B), int(L)),
SeqOffsets: []int32{position}, SeqOffsets: []int32{anchor},
SeqQueryLens: []int32{L}, SeqQueryLens: []int32{L},
} }
sliding, full := m.sharedHistories(b, caches) sliding, full := m.sharedHistories(ab, caches)
h := m.PreProjection.Forward(inputsEmbeds) h := m.PreProjection.Forward(inputsEmbeds)
positions := mlx.FromValues([]int32{position}, 1) positions := mlx.FromValues([]int32{anchor}, 1)
for _, layer := range m.Layers { for _, layer := range m.Layers {
h = layer.Forward(h, b, positions, B, L, &m.TextConfig, sliding, full) h = layer.Forward(h, ab, positions, B, L, &m.TextConfig, sliding, full)
} }
hidden = mlx.RMSNormFn(h, m.NormScaled, m.TextConfig.RMSNormEps) hidden = mlx.RMSNormFn(h, m.NormScaled, m.TextConfig.RMSNormEps)
projected := m.PostProjection.Forward(hidden) projected = m.PostProjection.Forward(hidden)
return m.unembed(hidden), projected return hidden, projected
} }
func (m *AssistantModel) sharedHistories(b *batch.Batch, caches []cache.Cache) (sliding, full *nn.KVHistory) { func (m *AssistantModel) sharedHistories(b *batch.Batch, caches []cache.Cache) (sliding, full *nn.KVHistory) {
@@ -302,7 +320,7 @@ func (m *AssistantModel) sharedHistories(b *batch.Batch, caches []cache.Cache) (
return sliding, full return sliding, full
} }
func (m *AssistantModel) unembed(hidden *mlx.Array) *mlx.Array { func (m *AssistantModel) Unembed(hidden *mlx.Array) *mlx.Array {
if m.UseOrderedEmbeddings { if m.UseOrderedEmbeddings {
return m.applyCentroidMasking(hidden) return m.applyCentroidMasking(hidden)
} }