ollama/x/create/cohere2moe.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

61 lines
2.4 KiB
Go

package create
import (
"encoding/json"
"fmt"
"strings"
)
// cohere2MoeImportTransform adjusts quantization for Cohere2 MoE imports
// (Command A family / North models).
type cohere2MoeImportTransform struct {
numLayers int
}
func newCohere2MoeImportTransform(rawConfig json.RawMessage) (quantizePolicy, error) {
var cfg struct {
NumHiddenLayers int `json:"num_hidden_layers"`
}
if err := json.Unmarshal(rawConfig, &cfg); err != nil {
return nil, fmt.Errorf("cohere2moe: parse config.json: %w", err)
}
return cohere2MoeImportTransform{numLayers: cfg.NumHiddenLayers}, nil
}
func (t cohere2MoeImportTransform) quantizationType(name string, shape []int32, quantize string) string {
base := normalizeQuantType(quantize)
// The embedding serves double duty: lookup (via QuantizedEmbedding) and the
// tied lm_head projection (via AsLinear). With a 262k vocab the bf16
// embedding dominates decode bandwidth through the lm_head matmul, so
// quantize it to the 8-bit variant of the requested mode, or keep source
// precision when that does not fit.
if isEmbedTokensWeight(name) && len(shape) == 2 {
return promoteEmbedding(shape, base)
}
// The MoE router picks the top-k expert set; quantization noise there can
// flip expert selection and compound downstream. It is tiny, so keep it in
// source precision. (GetTensorQuantization already skips "mlp.gate.weight";
// kept explicit here so renames in the default policy cannot regress this.)
if strings.HasSuffix(name, ".mlp.gate.weight") {
return ""
}
// Sensitive tensors (v_proj, k_proj, down_proj) get higher precision only
// at quantization-sensitive layer positions (useMoreBits) instead of the
// default policy's blanket promotion. The blanket int8 down_proj costs
// ~25% of decode bandwidth on a top-8 MoE; the layer-position heuristic
// keeps the early/late layers (and every third in between) at 8 bits where
// residual-stream error matters most.
isSensitive := strings.Contains(name, ".v_proj") || strings.Contains(name, ".k_proj") || strings.Contains(name, "down_proj")
if isSensitive && eightBit(base) != base && t.numLayers > 0 {
if idx := layerIndex(name); idx >= 0 {
// Bypass GetTensorQuantization's blanket promotion — the
// layer-position heuristic is authoritative here.
return sensitiveType(useMoreBits(idx, t.numLayers), shape, base)
}
}
return GetTensorQuantization(name, shape, quantize)
}