mlx: Gemma4 MTP speculative decoding (#15980)

This change adds support for MTP (multi-token prediction) speculative decoding for the
gemma4 model family.

It includes:
  * support for importing safetensors based gemma4 draft models with `ollama create`
  * a new DRAFT command in the Modelfile for specifying draft models
  * a --quantize-draft flag for the ollama create command to quantize the draft model
  * cache support for speculation
  * changes to the rotating cache to be able to handle MTP correctly
  * sampling support for draft model token prediction

---------

Co-authored-by: Daniel Hiltgen <daniel@ollama.com>
This commit is contained in:
Patrick Devine
2026-05-05 08:55:04 -07:00
committed by GitHub
parent 4017af96cd
commit 15e6076d79
28 changed files with 2928 additions and 42 deletions

View File

@@ -83,6 +83,8 @@ func (f Modelfile) CreateRequest(relativeDir string) (*api.CreateRequest, error)
req.Files[k] = v
}
}
case "draft":
return nil, errors.New("DRAFT requires --experimental")
case "adapter":
path, err := expandPath(c.Args, relativeDir)
if err != nil {
@@ -336,7 +338,7 @@ func (c Command) String() string {
switch c.Name {
case "model":
fmt.Fprintf(&sb, "FROM %s", c.Args)
case "license", "template", "system", "adapter", "renderer", "parser", "requires":
case "license", "template", "system", "adapter", "renderer", "parser", "requires", "draft":
fmt.Fprintf(&sb, "%s %s", strings.ToUpper(c.Name), quote(c.Args))
case "message":
role, message, _ := strings.Cut(c.Args, ": ")
@@ -362,7 +364,7 @@ const (
var (
errMissingFrom = errors.New("no FROM line")
errInvalidMessageRole = errors.New("message role must be one of \"system\", \"user\", or \"assistant\"")
errInvalidCommand = errors.New("command must be one of \"from\", \"license\", \"template\", \"system\", \"adapter\", \"renderer\", \"parser\", \"parameter\", \"message\", or \"requires\"")
errInvalidCommand = errors.New("command must be one of \"from\", \"license\", \"template\", \"system\", \"adapter\", \"draft\", \"renderer\", \"parser\", \"parameter\", \"message\", or \"requires\"")
)
type ParserError struct {
@@ -622,7 +624,7 @@ func isValidMessageRole(role string) bool {
func isValidCommand(cmd string) bool {
switch strings.ToLower(cmd) {
case "from", "license", "template", "system", "adapter", "renderer", "parser", "parameter", "message", "requires":
case "from", "license", "template", "system", "adapter", "draft", "renderer", "parser", "parameter", "message", "requires":
return true
default:
return false

View File

@@ -58,6 +58,32 @@ TEMPLATE """{{ if .System }}<|start_header_id|>system<|end_header_id|>
assert.Equal(t, expectedCommands, modelfile.Commands)
}
func TestParseFileDraft(t *testing.T) {
modelfile, err := ParseFile(strings.NewReader(`
FROM base
DRAFT ./assistant
`))
require.NoError(t, err)
expectedCommands := []Command{
{Name: "model", Args: "base"},
{Name: "draft", Args: "./assistant"},
}
assert.Equal(t, expectedCommands, modelfile.Commands)
assert.Contains(t, modelfile.String(), "DRAFT ./assistant")
}
func TestCreateRequestDraftRequiresExperimental(t *testing.T) {
modelfile, err := ParseFile(strings.NewReader(`
FROM base
DRAFT ./assistant
`))
require.NoError(t, err)
_, err = modelfile.CreateRequest("")
require.ErrorContains(t, err, "DRAFT requires --experimental")
}
func TestParseFileTrimSpace(t *testing.T) {
input := `
FROM " model 1"