diff --git a/x/mlxrunner/mtp.go b/x/mlxrunner/mtp.go index bc579ac52..e56d014ed 100644 --- a/x/mlxrunner/mtp.go +++ b/x/mlxrunner/mtp.go @@ -34,8 +34,6 @@ type mtpStats struct { accepted int mismatches int allAccepted int - batched int - serial int maxDraft int targetDuration time.Duration draftDuration time.Duration @@ -46,7 +44,6 @@ type mtpOptions struct { initialDraftTokens int maxDraftTokens int draftSchedule mtpDraftSchedule - serialValidate bool } func (r *Runner) mtpDefaults(sample bool) base.MTPDefaults { @@ -84,9 +81,6 @@ func (r *Runner) loadMTPOptions(sample bool) mtpOptions { if opts.initialDraftTokens > opts.maxDraftTokens { opts.initialDraftTokens = opts.maxDraftTokens } - if b, err := strconv.ParseBool(os.Getenv("OLLAMA_MLX_MTP_SERIAL_VALIDATE")); err == nil { - opts.serialValidate = b - } switch schedule := strings.ToLower(strings.TrimSpace(os.Getenv("OLLAMA_MLX_MTP_DRAFT_SCHEDULE"))); schedule { case "", string(mtpDraftScheduleConstant): opts.draftSchedule = mtpDraftScheduleConstant @@ -111,7 +105,9 @@ func positiveEnvInt(key string) int { return v } -func (r *Runner) useGreedyMTP(opts sampler.Options) bool { +// useMTP reports whether the request can run through the MTP speculative +// decode path. Logprobs are not yet supported. +func (r *Runner) useMTP(opts sampler.Options) bool { if r.Draft == nil { return false } @@ -121,190 +117,22 @@ func (r *Runner) useGreedyMTP(opts sampler.Options) bool { if _, ok := r.Model.(base.MTPEmbeddingModel); !ok { return false } - if !r.mtpDefaults(false).Enabled { + if !r.mtpDefaults(opts.Temperature != 0).Enabled { return false } if opts.Logprobs || opts.TopLogprobs > 0 { return false } - if opts.Temperature != 0 { - return false - } - repeatPenaltyNeutral := opts.RepeatPenalty <= 0 || opts.RepeatPenalty == 1 - topPNeutral := opts.TopP <= 0 || opts.TopP >= 1 - topKNeutral := opts.TopK <= 0 - return repeatPenaltyNeutral && opts.PresencePenalty == 0 && opts.FrequencyPenalty == 0 && topPNeutral && topKNeutral && opts.MinP == 0 + return true } -func (r *Runner) useSampleMTP(opts sampler.Options) bool { - if serial, err := strconv.ParseBool(os.Getenv("OLLAMA_MLX_MTP_SERIAL_VALIDATE")); err == nil && serial { - return false - } - if r.Draft == nil { - return false - } - if _, ok := r.Draft.(base.MTPDraftModel); !ok { - return false - } - if _, ok := r.Model.(base.MTPEmbeddingModel); !ok { - return false - } - if !r.mtpDefaults(true).Enabled { - return false - } - if opts.Logprobs || opts.TopLogprobs > 0 { - return false - } - return opts.Temperature != 0 -} - -func (r *Runner) runGreedyMTPDecode(ctx context.Context, request Request, session *cacheSession, caches []cache.Cache, seed []int32, position *int, started time.Time) error { +func (r *Runner) runMTPDecode(ctx context.Context, request Request, session *cacheSession, caches []cache.Cache, seed []int32, position *int, started time.Time) error { targetEmbeddings := r.Model.(base.MTPEmbeddingModel) draft := r.Draft.(base.MTPDraftModel) - mtpOpts := r.loadMTPOptions(false) + mtpOpts := r.loadMTPOptions(request.SamplerOpts.Temperature != 0) stats := mtpStats{maxDraft: mtpOpts.initialDraftTokens} draftLimit := mtpOpts.initialDraftTokens - slog.Info("MTP greedy decode enabled", "initial_draft_tokens", mtpOpts.initialDraftTokens, "max_draft_tokens", mtpOpts.maxDraftTokens, "draft_schedule", mtpOpts.draftSchedule, "serial_validate", mtpOpts.serialValidate) - - targetForward := func(token *mlx.Array) *mlx.Array { - fwd := r.Model.Forward(&batch.Batch{ - InputIDs: token, - SeqOffsets: []int32{int32(*position)}, - SeqQueryLens: []int32{int32(token.Dim(1))}, - }, caches) - *position += token.Dim(1) - return fwd - } - - hidden := targetForward(mlx.FromValues(seed, 1, len(seed))) - current := sampler.Result{Token: greedyTokenFromLogits(r.lastLogits(hidden))} - mlx.Pin(current.Arrays()...) - mlx.Sweep() - mlx.AsyncEval(current.Arrays()...) - defer func() { - mlx.Unpin(current.Arrays()...) - }() - - dec := decoder{tokenizer: r.Tokenizer} - final := CompletionResponse{Done: true, PromptEvalCount: len(request.Tokens), DoneReason: 1} - now := started - - generated := 0 - for generated < request.Options.NumPredict { - if err := ctx.Err(); err != nil { - return err - } - - t0 := time.Now() - hidden = targetForward(current.Token.ExpandDims(-1)) - baseLogits := r.lastLogits(hidden) - stats.targetDuration += time.Since(t0) - - if generated == 0 { - mlx.Eval(current.Arrays()...) - final.PromptEvalDuration = time.Since(now) - now = time.Now() - } - - done, err := r.emitTokens(ctx, request, session, &dec, []sampler.Result{current}, &final, &generated) - if err != nil { - return err - } - if done { - break - } - - stats.iterations++ - maxDraft := min(draftLimit, request.Options.NumPredict-generated) - t0 = time.Now() - draftTokens := r.generateMTPDrafts(draft, targetEmbeddings, current.Token, hidden, caches, int32(*position-1), maxDraft) - draftCount := 0 - if draftTokens != nil { - draftCount = draftTokens.Dim(1) - mlx.Pin(baseLogits, draftTokens) - mlx.Eval(draftTokens) - mlx.Sweep() - } - stats.draftDuration += time.Since(t0) - stats.drafted += draftCount - var next sampler.Result - if draftCount == 0 { - next = sampler.Result{Token: greedyTokenFromLogits(baseLogits)} - } else { - var accepted int - t0 = time.Now() - next, accepted, done, err = r.acceptMTPDrafts(ctx, request, session, &dec, caches, position, baseLogits, draftTokens, &final, &generated, &stats, mtpOpts) - stats.validateDuration += time.Since(t0) - mlx.Unpin(baseLogits, draftTokens) - if err != nil { - return err - } - stats.accepted += accepted - switch { - case mtpOpts.draftSchedule == mtpDraftScheduleConstant: - case accepted == draftCount: - stats.allAccepted++ - draftLimit = min(mtpOpts.maxDraftTokens, draftLimit+2) - default: - stats.mismatches++ - draftLimit = max(1, draftLimit-1) - } - if mtpOpts.draftSchedule == mtpDraftScheduleConstant { - if accepted == draftCount { - stats.allAccepted++ - } else { - stats.mismatches++ - } - } - stats.maxDraft = max(stats.maxDraft, draftLimit) - if next.Token == nil { - mlx.Sweep() - } - if done || generated >= request.Options.NumPredict { - break - } - } - - mlx.Pin(next.Arrays()...) - old := current - current = next - mlx.Unpin(old.Arrays()...) - mlx.Sweep() - mlx.AsyncEval(current.Arrays()...) - - if generated%256 == 0 { - mlx.ClearCache() - } - } - - final.EvalCount = generated - final.EvalDuration = time.Since(now) - acceptance := 0.0 - if stats.drafted > 0 { - acceptance = float64(stats.accepted) / float64(stats.drafted) - } - avgDraft := 0.0 - avgAccepted := 0.0 - if stats.iterations > 0 { - avgDraft = float64(stats.drafted) / float64(stats.iterations) - avgAccepted = float64(stats.accepted) / float64(stats.iterations) - } - slog.Info("MTP decode stats", "generated", generated, "drafted", stats.drafted, "accepted", stats.accepted, "acceptance", acceptance, "iterations", stats.iterations, "avg_draft", avgDraft, "avg_accepted", avgAccepted, "batched", stats.batched, "serial", stats.serial, "mismatches", stats.mismatches, "all_accepted", stats.allAccepted, "max_draft", stats.maxDraft, "draft_schedule", mtpOpts.draftSchedule, "target_duration", stats.targetDuration, "draft_duration", stats.draftDuration, "validate_duration", stats.validateDuration) - select { - case <-ctx.Done(): - return ctx.Err() - case request.Responses <- final: - return nil - } -} - -func (r *Runner) runSampleMTPDecode(ctx context.Context, request Request, session *cacheSession, caches []cache.Cache, seed []int32, position *int, started time.Time) error { - targetEmbeddings := r.Model.(base.MTPEmbeddingModel) - draft := r.Draft.(base.MTPDraftModel) - mtpOpts := r.loadMTPOptions(true) - stats := mtpStats{maxDraft: mtpOpts.initialDraftTokens} - draftLimit := mtpOpts.initialDraftTokens - slog.Info("MTP sample decode enabled", "initial_draft_tokens", mtpOpts.initialDraftTokens, "max_draft_tokens", mtpOpts.maxDraftTokens, "draft_schedule", mtpOpts.draftSchedule, "serial_validate", mtpOpts.serialValidate) + slog.Info("MTP decode enabled", "initial_draft_tokens", mtpOpts.initialDraftTokens, "max_draft_tokens", mtpOpts.maxDraftTokens, "draft_schedule", mtpOpts.draftSchedule) targetForward := func(token *mlx.Array) *mlx.Array { fwd := r.Model.Forward(&batch.Batch{ @@ -368,14 +196,13 @@ func (r *Runner) runSampleMTPDecode(ctx context.Context, request Request, sessio } stats.draftDuration += time.Since(t0) stats.drafted += draftCount - var next sampler.Result if draftCount == 0 { next = r.Sampler.Sample([]int{pipelineSlot}, baseLogits) } else { var accepted int t0 = time.Now() - next, accepted, done, err = r.acceptSampleMTPDrafts(ctx, request, session, &dec, caches, position, baseLogits, candidates, &final, &generated, &stats) + next, accepted, done, err = r.acceptMTPDrafts(ctx, request, session, &dec, caches, position, baseLogits, candidates, &final, &generated) stats.validateDuration += time.Since(t0) mlx.Unpin(candidateArrays...) if err != nil { @@ -431,7 +258,7 @@ func (r *Runner) runSampleMTPDecode(ctx context.Context, request Request, sessio avgDraft = float64(stats.drafted) / float64(stats.iterations) avgAccepted = float64(stats.accepted) / float64(stats.iterations) } - slog.Info("MTP decode stats", "mode", "sample", "generated", generated, "drafted", stats.drafted, "accepted", stats.accepted, "acceptance", acceptance, "iterations", stats.iterations, "avg_draft", avgDraft, "avg_accepted", avgAccepted, "batched", stats.batched, "serial", stats.serial, "mismatches", stats.mismatches, "all_accepted", stats.allAccepted, "max_draft", stats.maxDraft, "draft_schedule", mtpOpts.draftSchedule, "target_duration", stats.targetDuration, "draft_duration", stats.draftDuration, "validate_duration", stats.validateDuration) + slog.Info("MTP decode stats", "generated", generated, "drafted", stats.drafted, "accepted", stats.accepted, "acceptance", acceptance, "iterations", stats.iterations, "avg_draft", avgDraft, "avg_accepted", avgAccepted, "mismatches", stats.mismatches, "all_accepted", stats.allAccepted, "max_draft", stats.maxDraft, "draft_schedule", mtpOpts.draftSchedule, "target_duration", stats.targetDuration, "draft_duration", stats.draftDuration, "validate_duration", stats.validateDuration) select { case <-ctx.Done(): return ctx.Err() @@ -453,35 +280,6 @@ func (c *mtpDraftCandidates) Arrays() []*mlx.Array { return append([]*mlx.Array{c.tokens}, c.dist.Arrays()...) } -func (r *Runner) generateMTPDrafts(draft base.MTPDraftModel, target base.MTPEmbeddingModel, token *mlx.Array, hidden *mlx.Array, caches []cache.Cache, position int32, maxDraft int) *mlx.Array { - if maxDraft <= 0 { - return nil - } - - lastToken := token.ExpandDims(-1) - lastHidden := hidden - draftTokens := make([]*mlx.Array, 0, maxDraft) - - // Gemma4 assistant MTP is trained as "single-position" drafting: - // keep the RoPE/cache position anchored at the last target-seen token - // while the proposed token and projected hidden state advance. - for range maxDraft { - tokenEmbedding := target.TokenEmbeddings(lastToken) - inputs := tokenEmbedding.Concatenate(-1, lastHidden) - logits, projected := draft.Draft(inputs, position, caches) - stepLogits := r.lastLogitsFromLogits(logits) - nextToken := greedyTokenFromLogits(stepLogits) - - lastToken = nextToken.ExpandDims(-1) - lastHidden = projected - draftTokens = append(draftTokens, lastToken) - } - if len(draftTokens) == 0 { - return nil - } - return mlx.Concatenate(draftTokens, 1) -} - func (r *Runner) generateMTPDraftCandidates(draft base.MTPDraftModel, target base.MTPEmbeddingModel, token *mlx.Array, hidden *mlx.Array, caches []cache.Cache, position int32, maxDraft int) *mtpDraftCandidates { if maxDraft <= 0 { return nil @@ -523,16 +321,6 @@ func (r *Runner) generateMTPDraftCandidates(draft base.MTPDraftModel, target bas } } -func (r *Runner) acceptMTPDrafts(ctx context.Context, request Request, session *cacheSession, dec *decoder, caches []cache.Cache, position *int, baseLogits *mlx.Array, draftTokens *mlx.Array, final *CompletionResponse, generated *int, stats *mtpStats, opts mtpOptions) (sampler.Result, int, bool, error) { - if opts.serialValidate { - stats.serial++ - return r.acceptMTPDraftsSerial(ctx, request, session, dec, caches, position, baseLogits, draftTokens, final, generated) - } - - stats.batched++ - return r.acceptMTPDraftsBatched(ctx, request, session, dec, caches, position, baseLogits, draftTokens, final, generated) -} - // scheduleSpeculation schedules per-token snapshots at offsets // [before, before+draftCount) on every cache, so the speculative forward // captures a rollback point before each draft token's write. @@ -586,64 +374,10 @@ func commitSpeculation(caches []cache.Cache, accepted, draftCount, before int) { } } -func (r *Runner) acceptMTPDraftsBatched(ctx context.Context, request Request, session *cacheSession, dec *decoder, caches []cache.Cache, position *int, baseLogits *mlx.Array, draftTokens *mlx.Array, final *CompletionResponse, generated *int) (sampler.Result, int, bool, error) { - before := *position - draftCount := draftTokens.Dim(1) - - scheduleSpeculation(caches, before, draftCount) - hiddenSeq := r.Model.Forward(&batch.Batch{ - InputIDs: draftTokens, - SeqOffsets: []int32{int32(before)}, - SeqQueryLens: []int32{int32(draftCount)}, - }, caches) - - accepted := 0 - var next sampler.Result - done := false - - selectedTokens := r.mtpValidationTokens(baseLogits, hiddenSeq) - mlx.Eval(draftTokens, selectedTokens) - draftIDs := draftTokens.Ints() - selectedIDs := selectedTokens.Ints() - if len(selectedIDs) < draftCount+1 { - // Drain the scheduled snapshots and roll the speculative forward back - // out of the live caches before bailing, so the abandoned drafts don't - // reach the trie via session.close(). - commitSpeculation(caches, 0, draftCount, before) - return sampler.Result{}, accepted, false, fmt.Errorf("mtp validation produced %d tokens for %d draft tokens", len(selectedIDs), draftCount) - } - - for i, id := range draftIDs { - if selectedIDs[i] != id { - next = sampler.Result{Token: mtpTokenAt(selectedTokens, i)} - break - } - accepted++ - if r.Tokenizer.IsEOS(int32(id)) { - done = true - break - } - } - - commitSpeculation(caches, accepted, draftCount, before) - *position = before + accepted - - emitted, err := r.emitTokens(ctx, request, session, dec, draftResults(draftIDs[:accepted]), final, generated) - if err != nil { - return sampler.Result{}, accepted, emitted || done, err - } - if emitted || done { - return sampler.Result{}, accepted, true, nil - } - if next.Token == nil { - next = sampler.Result{Token: mtpTokenAt(selectedTokens, draftCount)} - } - return next, accepted, false, nil -} - -func (r *Runner) acceptSampleMTPDrafts(ctx context.Context, request Request, session *cacheSession, dec *decoder, caches []cache.Cache, position *int, baseLogits *mlx.Array, candidates *mtpDraftCandidates, final *CompletionResponse, generated *int, stats *mtpStats) (sampler.Result, int, bool, error) { - stats.batched++ - +// acceptMTPDrafts accepts the longest draft prefix that survives rejection +// sampling against the target model. At temperature 0 the distributions are +// point masses, so acceptance reduces to argmax-match. +func (r *Runner) acceptMTPDrafts(ctx context.Context, request Request, session *cacheSession, dec *decoder, caches []cache.Cache, position *int, baseLogits *mlx.Array, candidates *mtpDraftCandidates, final *CompletionResponse, generated *int) (sampler.Result, int, bool, error) { before := *position draftCount := candidates.tokens.Dim(1) scheduleSpeculation(caches, before, draftCount) @@ -753,42 +487,6 @@ func mtpTokenVector(token *mlx.Array) *mlx.Array { } } -func (r *Runner) acceptMTPDraftsSerial(ctx context.Context, request Request, session *cacheSession, dec *decoder, caches []cache.Cache, position *int, baseLogits *mlx.Array, draftTokens *mlx.Array, final *CompletionResponse, generated *int) (sampler.Result, int, bool, error) { - logits := baseLogits - accepted := 0 - draftIDs := draftTokens.Ints() - - for _, id := range draftIDs { - selected := greedyTokenFromLogits(logits) - mlx.Eval(selected) - selectedID := tokenID(selected) - if selectedID != id { - return sampler.Result{Token: mlx.FromValues([]int32{int32(selectedID)}, 1)}, accepted, false, nil - } - - hidden := r.Model.Forward(&batch.Batch{ - InputIDs: mlx.FromValues([]int32{int32(id)}, 1, 1), - SeqOffsets: []int32{int32(*position)}, - SeqQueryLens: []int32{1}, - }, caches) - (*position)++ - accepted++ - - res := sampler.Result{Token: mlx.FromValues([]int32{int32(id)}, 1)} - done, err := r.emitTokens(ctx, request, session, dec, []sampler.Result{res}, final, generated) - if err != nil { - return sampler.Result{}, accepted, done, err - } - if done { - return sampler.Result{}, accepted, true, nil - } - - logits = r.lastLogits(hidden) - } - - return sampler.Result{Token: greedyTokenFromLogits(logits)}, accepted, false, nil -} - // emitTokens records a run of generated tokens to session.outputs, then streams // them. A trailing EOS stops generation and is recorded but not streamed. // Returns whether to stop and any cancellation error. @@ -841,36 +539,21 @@ func (r *Runner) lastLogits(hidden *mlx.Array) *mlx.Array { return r.lastLogitsFromLogits(logits) } -func (r *Runner) mtpValidationTokens(baseLogits, hiddenSeq *mlx.Array) *mlx.Array { - return greedyTokenFromLogits(r.mtpValidationLogits(baseLogits, hiddenSeq)) -} - func (r *Runner) mtpValidationLogits(baseLogits, hiddenSeq *mlx.Array) *mlx.Array { seqLogits := r.Model.Unembed(hiddenSeq) return baseLogits.ExpandDims(1).Concatenate(1, seqLogits) } -func mtpTokenAt(tokens *mlx.Array, index int) *mlx.Array { - return tokens.Slice(mlx.Slice(), mlx.Slice(index)).Squeeze(0) -} - func (r *Runner) lastLogitsFromLogits(logits *mlx.Array) *mlx.Array { return logits.Slice(mlx.Slice(), mlx.Slice(logits.Dim(1)-1), mlx.Slice()).Squeeze(1) } -func greedyTokenFromLogits(logits *mlx.Array) *mlx.Array { - return logits.Argmax(-1, false).AsType(mlx.DTypeInt32) -} - +// tokenID reads a single-token array as its host id. It goes through the +// item accessor, which evaluates the array first: raw data reads on a lazy +// array race its evaluation and return garbage. func tokenID(token *mlx.Array) int { if token == nil { return -1 } - if token.DType() == mlx.DTypeInt32 { - ids := token.Ints() - if len(ids) > 0 { - return ids[0] - } - } return token.Int() } diff --git a/x/mlxrunner/mtp_test.go b/x/mlxrunner/mtp_test.go new file mode 100644 index 000000000..59cbf235e --- /dev/null +++ b/x/mlxrunner/mtp_test.go @@ -0,0 +1,446 @@ +package mlxrunner + +import ( + "context" + "encoding/json" + "fmt" + "slices" + "strings" + "testing" + "time" + + "github.com/ollama/ollama/api" + "github.com/ollama/ollama/x/mlxrunner/batch" + "github.com/ollama/ollama/x/mlxrunner/cache" + "github.com/ollama/ollama/x/mlxrunner/mlx" + "github.com/ollama/ollama/x/mlxrunner/model/base" + sampler "github.com/ollama/ollama/x/mlxrunner/sample" + "github.com/ollama/ollama/x/tokenizer" +) + +// skipIfNoMLX skips a test that exercises native MLX when the dynamic library +// is unavailable, as on CI runners without an MLX build. +func skipIfNoMLX(t *testing.T) { + t.Helper() + if err := mlx.CheckInit(); err != nil { + t.Skipf("MLX not available: %v", err) + } +} + +// The MTP fakes make hidden state and logits the same tensor (Forward returns +// one-hot logits, Unembed is the identity), so tests fully script target and +// draft predictions. + +const mtpTestVocab = 8 + +// oneHotLogits builds logits with a large value on each listed token id. +func oneHotLogits(tokens []int32) *mlx.Array { + data := make([]float32, len(tokens)*mtpTestVocab) + for i, tok := range tokens { + data[i*mtpTestVocab+int(tok)] = 30 + } + return mlx.FromValues(data, 1, len(tokens), mtpTestVocab) +} + +// fakeMTPModel is a target whose next-token prediction is a fixed function of +// the input token; Forward also feeds ids to the caches so offsets advance. +type fakeMTPModel struct { + predict map[int32]int32 + tok *tokenizer.Tokenizer + // forwards records each Forward call so tests can assert contiguous writes. + forwards []forwardCall +} + +type forwardCall struct { + offset int32 + n int32 +} + +func (m *fakeMTPModel) Forward(b *batch.Batch, caches []cache.Cache) *mlx.Array { + mlx.Eval(b.InputIDs) + ids := b.InputIDs.Ints() + m.forwards = append(m.forwards, forwardCall{offset: b.SeqOffsets[0], n: int32(len(ids))}) + for _, c := range caches { + if rc, ok := c.(*fakeRewindableCache); ok { + seg := make([]int32, len(ids)) + for i, id := range ids { + seg[i] = int32(id) + } + rc.feed(seg) + } + } + + preds := make([]int32, len(ids)) + for i, id := range ids { + preds[i] = m.predict[int32(id)] + } + return oneHotLogits(preds) +} + +func (m *fakeMTPModel) Unembed(x *mlx.Array) *mlx.Array { return x } +func (m *fakeMTPModel) NumLayers() int { return 1 } +func (m *fakeMTPModel) Tokenizer() *tokenizer.Tokenizer { return m.tok } +func (m *fakeMTPModel) MaxContextLength() int { return 4096 } +func (m *fakeMTPModel) LoadWeights(map[string]*mlx.Array) error { + return nil +} + +// TokenEmbeddings returns a width-1 embedding holding the token id as a float, +// 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 ( + _ base.Model = (*fakeMTPModel)(nil) + _ 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 { + predict map[int32]int32 + // calls records each Draft call so tests can assert the position convention. + calls []draftCall +} + +type draftCall struct { + position int32 + from int32 +} + +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) { + mlx.Eval(inputEmbeds) + prev := int32(inputEmbeds.Floats()[0]) + d.calls = append(d.calls, draftCall{position: position, from: prev}) + return oneHotLogits([]int32{d.predict[prev]}), mlx.Zeros(mlx.DTypeFloat32, 1, 1, mtpTestVocab) +} + +var ( + _ base.DraftModel = (*fakeMTPDraft)(nil) + _ base.MTPDraftModel = (*fakeMTPDraft)(nil) +) + +// 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 +// IsEOS reports membership. +func newTestTokenizer(t *testing.T, eos []int32) *tokenizer.Tokenizer { + t.Helper() + vocab := make(map[string]int32, mtpTestVocab) + for i := range mtpTestVocab { + vocab[fmt.Sprintf("%d", i)] = int32(i) + } + model := map[string]any{ + "type": "BPE", + "vocab": vocab, + "merges": []string{}, + } + data, err := json.Marshal(map[string]any{"model": model}) + if err != nil { + t.Fatalf("marshal tokenizer: %v", err) + } + genConfig, err := json.Marshal(map[string]any{"eos_token_id": eos}) + if err != nil { + t.Fatalf("marshal generation config: %v", err) + } + tok, err := tokenizer.LoadFromBytesWithConfig(data, &tokenizer.TokenizerConfig{GenerationConfigJSON: genConfig}) + if err != nil { + t.Fatalf("load tokenizer: %v", err) + } + return tok +} + +// mtpTestRunner wires a Runner with the MTP fakes and a real sampler +// registered with opts. +func mtpTestRunner(t *testing.T, predict map[int32]int32, eos []int32, opts sampler.Options) *Runner { + t.Helper() + tok := newTestTokenizer(t, eos) + r := &Runner{ + Model: &fakeMTPModel{predict: predict, tok: tok}, + Tokenizer: tok, + Sampler: sampler.New(4096), + } + r.Sampler.Add(pipelineSlot, opts, nil) + t.Cleanup(func() { r.Sampler.Remove(pipelineSlot) }) + return r +} + +// collectResponses drains a buffered response channel into the concatenated +// streamed content and the captured terminal response. +func collectResponses(ch chan CompletionResponse) (content string, final CompletionResponse) { + var b strings.Builder + for { + select { + case resp := <-ch: + if resp.Done { + return b.String(), resp + } + b.WriteString(resp.Content) + default: + return b.String(), final + } + } +} + +func TestAcceptMTPDraftsGreedyAcceptAll(t *testing.T) { + skipIfNoMLX(t) + // Target predicts 1->2->3->4 along the accepted chain; the draft proposed + // exactly that, so every draft token is accepted and the bonus token is the + // target's prediction after the last accepted token. + predict := map[int32]int32{1: 2, 2: 3, 3: 4, 4: 5} + r := mtpTestRunner(t, predict, nil, sampler.Options{}) + + caches, _ := newMTPTestCaches(1) + candidates := scriptedCandidates(r, []int32{2, 3, 4}) + baseLogits := oneHotLogits([]int32{2}).Squeeze(1) // target prediction after the seed token 1 + + session, ch := newMTPTestSession(caches) + position := caches[0].Offset() + final := CompletionResponse{Done: true} + generated := 0 + + req := Request{Responses: ch, CompletionRequest: CompletionRequest{Options: api.Options{NumPredict: 100}}} + next, accepted, done, err := r.acceptMTPDrafts(context.Background(), req, session, &decoder{tokenizer: r.Tokenizer}, caches, &position, baseLogits, candidates, &final, &generated) + if err != nil { + t.Fatalf("acceptMTPDrafts: %v", err) + } + if accepted != 3 { + t.Fatalf("accepted = %d, want 3", accepted) + } + if done { + t.Fatalf("done = true, want false") + } + if got := tokenID(next.Token); got != 5 { + t.Fatalf("bonus token = %d, want 5", got) + } + if position != 3 { + t.Fatalf("position = %d, want 3", position) + } + if got := caches[0].Offset(); got != 3 { + t.Fatalf("cache offset = %d, want 3 (all drafts kept)", got) + } + if generated != 3 { + t.Fatalf("generated = %d, want 3", generated) + } +} + +func TestAcceptMTPDraftsGreedyMismatch(t *testing.T) { + skipIfNoMLX(t) + // Target predicts 1->2->9 but the draft proposed 2 then 7: the second draft + // token mismatches, so only the first is accepted and the bonus is the + // target's own prediction (3) at the rejection point. + predict := map[int32]int32{1: 2, 2: 3, 7: 0} + r := mtpTestRunner(t, predict, nil, sampler.Options{}) + + caches, _ := newMTPTestCaches(1) + candidates := scriptedCandidates(r, []int32{2, 7}) + baseLogits := oneHotLogits([]int32{2}).Squeeze(1) + + session, ch := newMTPTestSession(caches) + position := caches[0].Offset() + final := CompletionResponse{Done: true} + generated := 0 + + req := Request{Responses: ch, CompletionRequest: CompletionRequest{Options: api.Options{NumPredict: 100}}} + next, accepted, done, err := r.acceptMTPDrafts(context.Background(), req, session, &decoder{tokenizer: r.Tokenizer}, caches, &position, baseLogits, candidates, &final, &generated) + if err != nil { + t.Fatalf("acceptMTPDrafts: %v", err) + } + if accepted != 1 { + t.Fatalf("accepted = %d, want 1", accepted) + } + if done { + t.Fatalf("done = true, want false") + } + if got := tokenID(next.Token); got != 3 { + t.Fatalf("bonus token = %d, want 3 (target prediction at rejection)", got) + } + if position != 1 { + t.Fatalf("position = %d, want 1", position) + } + if got := caches[0].Offset(); got != 1 { + t.Fatalf("cache offset = %d, want 1 (rolled back to accepted)", got) + } +} + +func TestAcceptMTPDraftsGreedyEOS(t *testing.T) { + skipIfNoMLX(t) + // The second accepted draft token is EOS: it is recorded but stops + // generation, no bonus token is produced, and the cache keeps both tokens. + const eos int32 = 6 + predict := map[int32]int32{1: 2, 2: eos, eos: 0} + r := mtpTestRunner(t, predict, []int32{eos}, sampler.Options{}) + + caches, _ := newMTPTestCaches(1) + candidates := scriptedCandidates(r, []int32{2, eos}) + baseLogits := oneHotLogits([]int32{2}).Squeeze(1) + + session, ch := newMTPTestSession(caches) + position := caches[0].Offset() + final := CompletionResponse{Done: true, DoneReason: 1} + generated := 0 + + req := Request{Responses: ch, CompletionRequest: CompletionRequest{Options: api.Options{NumPredict: 100}}} + next, accepted, done, err := r.acceptMTPDrafts(context.Background(), req, session, &decoder{tokenizer: r.Tokenizer}, caches, &position, baseLogits, candidates, &final, &generated) + if err != nil { + t.Fatalf("acceptMTPDrafts: %v", err) + } + if accepted != 2 { + t.Fatalf("accepted = %d, want 2 (token + EOS)", accepted) + } + if !done { + t.Fatalf("done = false, want true") + } + if next.Token != nil { + t.Fatalf("bonus token = %d, want none after EOS", tokenID(next.Token)) + } + if final.DoneReason != 0 { + t.Fatalf("DoneReason = %d, want 0 (EOS)", final.DoneReason) + } + if position != 2 { + t.Fatalf("position = %d, want 2", position) + } +} + +func TestRunMTPDecodeGreedy(t *testing.T) { + skipIfNoMLX(t) + // The seed token 1 is the last prefill token; its prediction (2) is the + // first generated token. The decode then walks 2->3->4->EOS. The draft + // proposes the correct chain so steps are accepted in a single forward. + const eos int32 = 7 + predict := map[int32]int32{1: 2, 2: 3, 3: 4, 4: eos, eos: 0} + r := mtpTestRunner(t, predict, []int32{eos}, sampler.Options{}) + // The draft mirrors the target chain, so every drafted token is accepted. + draft := &fakeMTPDraft{predict: predict} + r.Draft = draft + + caches, _ := newMTPTestCaches(1) + session, ch := newMTPTestSession(caches) + position := 1 // one prefill token already processed + + req := Request{ + Responses: ch, + Tokens: []int32{0}, + CompletionRequest: CompletionRequest{Options: api.Options{NumPredict: 20}}, + SamplerOpts: sampler.Options{}, + } + if err := r.runMTPDecode(context.Background(), req, session, caches, []int32{1}, &position, time.Now()); err != nil { + t.Fatalf("runMTPDecode: %v", err) + } + + content, final := collectResponses(ch) + if content != "234" { + t.Fatalf("content = %q, want %q", content, "234") + } + if !final.Done { + t.Fatalf("final response not marked Done") + } + if final.DoneReason != 0 { + t.Fatalf("DoneReason = %d, want 0 (EOS)", final.DoneReason) + } + if got := []int32{2, 3, 4, eos}; !slices.Equal(session.outputs, got) { + t.Fatalf("session outputs = %v, want %v", session.outputs, got) + } + + // The target writes the caches contiguously: the seed token at offset 1, + // the round's current token at offset 2, then the 4-token validation + // forward at offset 3. + wantForwards := []forwardCall{{offset: 1, n: 1}, {offset: 2, n: 1}, {offset: 3, n: 4}} + model := r.Model.(*fakeMTPModel) + if !slices.Equal(model.forwards, wantForwards) { + t.Fatalf("target forwards = %v, want %v", model.forwards, wantForwards) + } + + // Single-position drafting anchors every Draft call in a round at the + // last target-seen position (the current token, offset 2), while the + // extended token advances along the proposed chain. + wantDraft := []draftCall{{2, 2}, {2, 3}, {2, 4}, {2, eos}} + if !slices.Equal(draft.calls, wantDraft) { + t.Fatalf("draft calls = %v, want %v", draft.calls, wantDraft) + } +} + +func TestRunMTPDecodeSampled(t *testing.T) { + skipIfNoMLX(t) + // The same chain at temperature 1: because oneHotLogits uses a large gap, + // the proposal and target distributions are effectively point masses, so the + // rejection-sampling accept path that the sampled and greedy paths now share + // accepts the mirrored draft chain deterministically. + const eos int32 = 7 + predict := map[int32]int32{1: 2, 2: 3, 3: 4, 4: eos, eos: 0} + r := mtpTestRunner(t, predict, []int32{eos}, sampler.Options{Temperature: 1, Seed: 42, UseSeed: true}) + r.Draft = &fakeMTPDraft{predict: predict} + + if !r.useMTP(sampler.Options{Temperature: 1}) { + t.Fatalf("useMTP rejected a sampled request") + } + + caches, _ := newMTPTestCaches(1) + session, ch := newMTPTestSession(caches) + position := 1 + + req := Request{ + Responses: ch, + Tokens: []int32{0}, + CompletionRequest: CompletionRequest{Options: api.Options{NumPredict: 20}}, + SamplerOpts: sampler.Options{Temperature: 1, Seed: 42, UseSeed: true}, + } + if err := r.runMTPDecode(context.Background(), req, session, caches, []int32{1}, &position, time.Now()); err != nil { + t.Fatalf("runMTPDecode: %v", err) + } + + content, final := collectResponses(ch) + if content != "234" { + t.Fatalf("content = %q, want %q", content, "234") + } + if final.DoneReason != 0 { + t.Fatalf("DoneReason = %d, want 0 (EOS)", final.DoneReason) + } + if got := []int32{2, 3, 4, eos}; !slices.Equal(session.outputs, got) { + t.Fatalf("session outputs = %v, want %v", session.outputs, got) + } +} + +// newMTPTestCaches returns n rewindable fake caches sharing one snapshot +// tracker, matching the cache.Cache the speculation helpers drive. +func newMTPTestCaches(n int) ([]cache.Cache, *snapshotTracker) { + tr := &snapshotTracker{} + caches := make([]cache.Cache, n) + for i := range caches { + caches[i] = &fakeRewindableCache{tracker: tr} + } + return caches, tr +} + +// newMTPTestSession wraps caches in a cacheSession with a buffered response +// channel large enough to hold a short decode run without a reader. +func newMTPTestSession(caches []cache.Cache) (*cacheSession, chan CompletionResponse) { + ch := make(chan CompletionResponse, 256) + return &cacheSession{caches: caches}, ch +} + +// scriptedCandidates builds draft candidates by running the real generator +// against a draft whose prediction chain, starting from seed token 0, yields +// exactly the requested tokens. Using the real generator means the proposal +// distributions match what acceptMTPDrafts expects. +func scriptedCandidates(r *Runner, tokens []int32) *mtpDraftCandidates { + chain := map[int32]int32{} + prev := int32(0) + for _, tok := range tokens { + chain[prev] = tok + prev = tok + } + draft := &fakeMTPDraft{predict: chain} + target := r.Model.(base.MTPEmbeddingModel) + seed := mlx.FromValues([]int32{0}, 1, 1) + hidden := mlx.Zeros(mlx.DTypeFloat32, 1, 1, mtpTestVocab) + return r.generateMTPDraftCandidates(draft, target, seed, hidden, nil, 0, len(tokens)) +} diff --git a/x/mlxrunner/pipeline.go b/x/mlxrunner/pipeline.go index 0f69374e6..e57201354 100644 --- a/x/mlxrunner/pipeline.go +++ b/x/mlxrunner/pipeline.go @@ -137,11 +137,8 @@ func (r *Runner) TextGenerationPipeline(ctx context.Context, request Request) er // Register the sampler after prefill completes. r.Sampler.Add(pipelineSlot, request.SamplerOpts, inputs) - if r.useGreedyMTP(request.SamplerOpts) { - return r.runGreedyMTPDecode(ctx, request, session, caches, tokens[processed:], &position, now) - } - if r.useSampleMTP(request.SamplerOpts) { - return r.runSampleMTPDecode(ctx, request, session, caches, tokens[processed:], &position, now) + if r.useMTP(request.SamplerOpts) { + return r.runMTPDecode(ctx, request, session, caches, tokens[processed:], &position, now) } step := func(token *mlx.Array) sampler.Result {