ollama/x/create/pipeline.go
Patrick Devine 964ea42c09
mlx: x/create rewrite (#16919)
This is a rewrite of the create functionality for the MLX engine.

The core idea behind the create functionality is to break the import/convert into a pipeline of distinct phases:

* Read (scan the safetensors directory for the various bits of metadata)
* Classify (determine what the import type)
* Plan (determine any transforms that need to be done)
* Write (transform any data as necessary and write out the blobs)
* Create the manifest

Each architecture has a "policy" which determines how to convert the model correctly. A number of different formats for safetensors are supported including:

* nvfp4 (two formats: model optimized, torch)
* fp8 datatypes (convert to mxfp8)
* standard bf16 based weights

A number of cleanups/simplifications have been done including:

* using the baked in names for the tensors instead of munging them into something else
* unified 3d expert tensors (instead of separate per expert tensors)
* fewer unnecessary transforms to the various tensors in a model (keep a model as close to the source as possible)
* unified capability checking
* draft model handling (for MTP) is done on the same path

Image generation has been intentionally removed.
2026-07-03 18:30:45 -07:00

121 lines
3.8 KiB
Go

package create
import (
"fmt"
"io"
"os"
"path/filepath"
"strings"
)
// Create imports a safetensors model through the full pipeline: read the
// source into an inventory, classify it, plan the output blobs, write them
// through store, import the config files, and write the manifest. It is the
// server-side entry point — the caller supplies blob storage (store) and
// manifest assembly (writeManifest).
func Create(modelName, modelDir, quantize string, store BlobStore, writeManifest ManifestWriter, fn func(status string)) error {
defer sweepMLX()
inv, err := ReadInventory(modelDir)
if err != nil {
return fmt.Errorf("read model: %w", err)
}
class, err := Classify(inv, quantize)
if err != nil {
return err
}
policy, err := newTensorImportTransform(inv)
if err != nil {
return fmt.Errorf("build quantization policy for %q: %w", inv.Config.Architecture(), err)
}
specs, err := Plan(inv, class, policy)
if err != nil {
return fmt.Errorf("plan model: %w", err)
}
fn(fmt.Sprintf("importing %s (%d tensors%s)", modelName, len(inv.Tensors), quantizeStatus(class)))
layers, err := WriteBlobs(specs, modelDir, store)
if err != nil {
return err
}
// Import config files (config.json, tokenizer, etc.) as JSON blobs.
configLayers, configLayer, err := importConfigBlobs(modelDir, "", store, fn)
if err != nil {
return err
}
layers = append(layers, configLayers...)
if configLayer.Digest == "" {
return fmt.Errorf("config.json not found in %s", modelDir)
}
fn(fmt.Sprintf("writing manifest for %s", modelName))
if err := writeManifest(modelName, configLayer, layers); err != nil {
return fmt.Errorf("write manifest: %w", err)
}
fn(fmt.Sprintf("successfully imported %s with %d layers", modelName, len(layers)))
return nil
}
const mediaTypeImageJSON = "application/vnd.ollama.image.json"
// importConfigBlobs writes every .json in modelDir (except the shard index) as an
// image.json blob, prefixing each blob name with namePrefix, and returns the
// resulting layers along with the config.json layer (zero value if absent). The
// target import passes "" for namePrefix; a draft import passes "draft/" so its
// config sits beside the target's.
func importConfigBlobs(modelDir, namePrefix string, store BlobStore, fn func(status string)) ([]LayerInfo, LayerInfo, error) {
entries, err := os.ReadDir(modelDir)
if err != nil {
return nil, LayerInfo{}, err
}
var layers []LayerInfo
var configLayer LayerInfo
for _, entry := range entries {
if entry.IsDir() || !strings.HasSuffix(entry.Name(), ".json") || entry.Name() == "model.safetensors.index.json" {
continue
}
name := entry.Name()
fn(fmt.Sprintf("importing config %s", name))
f, err := os.Open(filepath.Join(modelDir, name))
if err != nil {
return nil, LayerInfo{}, fmt.Errorf("open %s: %w", name, err)
}
layer, err := store.WriteBlob(f, mediaTypeImageJSON, namePrefix+name)
f.Close()
if err != nil {
return nil, LayerInfo{}, fmt.Errorf("write config %s: %w", name, err)
}
if name == "config.json" {
configLayer = layer
}
layers = append(layers, layer)
}
return layers, configLayer, nil
}
func quantizeStatus(c Classification) string {
switch c.Kind {
case SourceBlockFP8:
return ", converting fp8 to mxfp8"
case SourcePrequantized:
return ", preserving source quantization"
default:
if c.Quantize != "" {
return ", quantizing to " + c.Quantize
}
return ""
}
}
// StoreFromLayerCreator adapts a LayerCreator-style function to a BlobStore, so
// a caller that already has a blob-writing callback can drive the pipeline.
func StoreFromLayerCreator(fn LayerCreator) BlobStore {
return layerCreatorStore{fn}
}
type layerCreatorStore struct{ fn LayerCreator }
func (s layerCreatorStore) WriteBlob(r io.Reader, mediaType, name string) (LayerInfo, error) {
return s.fn(r, mediaType, name)
}