mlx: fix reported information in ollama show (#16289)
This change updates the show API for MLX models to: * display the correct quantization in mixed precision models * not display the global_scale scalar value * not duplicate the `tools` capability
This commit is contained in:
+1
-1
@@ -132,7 +132,7 @@ func (m *Model) Capabilities() []model.Capability {
|
||||
if err != nil {
|
||||
slog.Warn("model template contains errors", "error", err)
|
||||
}
|
||||
if slices.Contains(v, "tools") || (builtinParser != nil && builtinParser.HasToolSupport()) {
|
||||
if !slices.Contains(capabilities, model.CapabilityTools) && (slices.Contains(v, "tools") || (builtinParser != nil && builtinParser.HasToolSupport())) {
|
||||
capabilities = append(capabilities, model.CapabilityTools)
|
||||
}
|
||||
|
||||
|
||||
@@ -138,6 +138,17 @@ func TestModelCapabilities(t *testing.T) {
|
||||
},
|
||||
expectedCaps: []model.Capability{model.CapabilityCompletion, model.CapabilityTools},
|
||||
},
|
||||
{
|
||||
name: "model with tools capability from config and parser",
|
||||
model: Model{
|
||||
Config: model.ConfigV2{
|
||||
Capabilities: []string{"completion", "tools"},
|
||||
Parser: "qwen3-coder",
|
||||
},
|
||||
Template: chatTemplate,
|
||||
},
|
||||
expectedCaps: []model.Capability{model.CapabilityCompletion, model.CapabilityTools},
|
||||
},
|
||||
{
|
||||
name: "model with vision capability",
|
||||
model: Model{
|
||||
|
||||
+37
-7
@@ -306,7 +306,7 @@ func getTensorInfoFromManifest(mf *manifest.Manifest) ([]api.Tensor, error) {
|
||||
}
|
||||
|
||||
// GetSafetensorsDtype returns the quantization type for a safetensors model.
|
||||
// Reads tensor headers until quantized weights are found.
|
||||
// Reads tensor headers and reports the lowest-precision quantized weight type.
|
||||
// Falls back to torch_dtype from config.json if no quant metadata exists.
|
||||
func GetSafetensorsDtype(name model.Name) (string, error) {
|
||||
mf, err := manifest.ParseNamedManifest(name)
|
||||
@@ -314,8 +314,9 @@ func GetSafetensorsDtype(name model.Name) (string, error) {
|
||||
return "", fmt.Errorf("failed to load manifest: %w", err)
|
||||
}
|
||||
|
||||
// Mixed models can start with unquantized embeddings or heads, so scan until
|
||||
// any tensor blob reports quantized weight metadata.
|
||||
var bestQuantType string
|
||||
var bestPrecision int
|
||||
|
||||
for _, layer := range mf.Layers {
|
||||
if layer.MediaType != manifest.MediaTypeImageTensor {
|
||||
continue
|
||||
@@ -334,12 +335,22 @@ func GetSafetensorsDtype(name model.Name) (string, error) {
|
||||
continue
|
||||
}
|
||||
for _, info := range infos {
|
||||
if quantType := canonicalQuantType(info.QuantType); quantType != "" {
|
||||
return quantType, nil
|
||||
quantType := canonicalQuantType(info.QuantType)
|
||||
if quantType == "" {
|
||||
continue
|
||||
}
|
||||
precision := quantTypePrecision(quantType)
|
||||
if bestQuantType == "" || (precision > 0 && (bestPrecision == 0 || precision < bestPrecision)) {
|
||||
bestQuantType = quantType
|
||||
bestPrecision = precision
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if bestQuantType != "" {
|
||||
return bestQuantType, nil
|
||||
}
|
||||
|
||||
// Not quantized - return torch_dtype from config.json
|
||||
var cfg struct {
|
||||
TorchDtype string `json:"torch_dtype"`
|
||||
@@ -351,6 +362,17 @@ func GetSafetensorsDtype(name model.Name) (string, error) {
|
||||
return cfg.TorchDtype, nil
|
||||
}
|
||||
|
||||
func quantTypePrecision(quantType string) int {
|
||||
switch canonicalQuantType(quantType) {
|
||||
case "int4", "nvfp4", "mxfp4":
|
||||
return 4
|
||||
case "int8", "mxfp8":
|
||||
return 8
|
||||
default:
|
||||
return 0
|
||||
}
|
||||
}
|
||||
|
||||
// safetensorsTensorInfo holds metadata about a tensor from a safetensors header
|
||||
type safetensorsTensorInfo struct {
|
||||
Name string // tensor name from the header key
|
||||
@@ -361,7 +383,8 @@ type safetensorsTensorInfo struct {
|
||||
}
|
||||
|
||||
// parseSafetensorsAllHeaders parses all tensor entries from a safetensors header.
|
||||
// Returns one safetensorsTensorInfo per main tensor (skipping __metadata__, .scale, .bias).
|
||||
// Returns one safetensorsTensorInfo per main tensor, skipping quantization
|
||||
// companion entries such as __metadata__, .scale, .bias, and .global_scale.
|
||||
// For packed blobs this returns multiple entries; for single-tensor blobs, one entry.
|
||||
// Each tensor's quant type is inferred from its shape and the presence of .scale/.bias entries
|
||||
// when no global __metadata__ quant_type is present.
|
||||
@@ -404,7 +427,7 @@ func parseSafetensorsAllHeaders(r io.Reader) ([]safetensorsTensorInfo, error) {
|
||||
// Collect all main tensor entries (sorted for deterministic output)
|
||||
var mainNames []string
|
||||
for name := range header {
|
||||
if name == "__metadata__" || strings.HasSuffix(name, ".scale") || strings.HasSuffix(name, ".bias") {
|
||||
if isSafetensorsCompanionTensor(name) {
|
||||
continue
|
||||
}
|
||||
mainNames = append(mainNames, name)
|
||||
@@ -438,6 +461,13 @@ func parseSafetensorsAllHeaders(r io.Reader) ([]safetensorsTensorInfo, error) {
|
||||
return results, nil
|
||||
}
|
||||
|
||||
func isSafetensorsCompanionTensor(name string) bool {
|
||||
return name == "__metadata__" ||
|
||||
strings.HasSuffix(name, ".scale") ||
|
||||
strings.HasSuffix(name, ".bias") ||
|
||||
strings.HasSuffix(name, ".global_scale")
|
||||
}
|
||||
|
||||
// inferQuantType infers the quantization type for a tensor from its shape and scale shape.
|
||||
// Returns "int4", "int8", etc. or "" if not quantized.
|
||||
func inferQuantType(header map[string]json.RawMessage, name string) string {
|
||||
|
||||
+50
-5
@@ -833,6 +833,34 @@ func TestParseSafetensorsAllHeaders(t *testing.T) {
|
||||
wantDtypes: []string{"U32", "U32"},
|
||||
wantQuants: []string{"int4", "int4"},
|
||||
},
|
||||
{
|
||||
name: "nvfp4 global scale is hidden",
|
||||
header: map[string]any{
|
||||
"__metadata__": map[string]any{
|
||||
"quant_type": "nvfp4",
|
||||
"group_size": "16",
|
||||
},
|
||||
"model.layers.0.mlp.up_proj.weight": map[string]any{
|
||||
"dtype": "U32",
|
||||
"shape": []int64{2560, 320},
|
||||
"data_offsets": []int64{0, 3276800},
|
||||
},
|
||||
"model.layers.0.mlp.up_proj.weight.scale": map[string]any{
|
||||
"dtype": "U8",
|
||||
"shape": []int64{2560, 160},
|
||||
"data_offsets": []int64{3276800, 3686400},
|
||||
},
|
||||
"model.layers.0.mlp.up_proj.weight.global_scale": map[string]any{
|
||||
"dtype": "F32",
|
||||
"shape": []int64{1},
|
||||
"data_offsets": []int64{3686400, 3686404},
|
||||
},
|
||||
},
|
||||
wantCount: 1,
|
||||
wantNames: []string{"model.layers.0.mlp.up_proj.weight"},
|
||||
wantDtypes: []string{"U32"},
|
||||
wantQuants: []string{"nvfp4"},
|
||||
},
|
||||
{
|
||||
name: "packed mixed-precision blob (no global metadata)",
|
||||
header: map[string]any{
|
||||
@@ -1049,7 +1077,7 @@ func TestGetTensorInfoFromManifest_Packed(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetSafetensorsDtypeScansPastUnquantizedFirstBlob(t *testing.T) {
|
||||
func TestGetSafetensorsDtypeChoosesLowestPrecisionQuantizedBlob(t *testing.T) {
|
||||
t.Setenv("OLLAMA_MODELS", t.TempDir())
|
||||
|
||||
writeSafetensorsLayer := func(t *testing.T, header map[string]any, name string) manifest.Layer {
|
||||
@@ -1110,8 +1138,25 @@ func TestGetSafetensorsDtypeScansPastUnquantizedFirstBlob(t *testing.T) {
|
||||
},
|
||||
}, "model.layers.0.mlp.down_proj.weight")
|
||||
|
||||
name := model.ParseName("mixed-fp8-safetensors")
|
||||
if err := manifest.WriteManifest(name, configLayer, []manifest.Layer{unquantized, quantized}); err != nil {
|
||||
lowerPrecision := writeSafetensorsLayer(t, map[string]any{
|
||||
"__metadata__": map[string]string{
|
||||
"quant_type": "nvfp4",
|
||||
"group_size": "16",
|
||||
},
|
||||
"model.layers.0.mlp.up_proj.weight": map[string]any{
|
||||
"dtype": "U32",
|
||||
"shape": []int64{16, 2},
|
||||
"data_offsets": []int64{0, 128},
|
||||
},
|
||||
"model.layers.0.mlp.up_proj.weight.scale": map[string]any{
|
||||
"dtype": "U8",
|
||||
"shape": []int64{16, 1},
|
||||
"data_offsets": []int64{128, 144},
|
||||
},
|
||||
}, "model.layers.0.mlp.up_proj.weight")
|
||||
|
||||
name := model.ParseName("mixed-precision-safetensors")
|
||||
if err := manifest.WriteManifest(name, configLayer, []manifest.Layer{unquantized, quantized, lowerPrecision}); err != nil {
|
||||
t.Fatalf("failed to write manifest: %v", err)
|
||||
}
|
||||
|
||||
@@ -1119,7 +1164,7 @@ func TestGetSafetensorsDtypeScansPastUnquantizedFirstBlob(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("GetSafetensorsDtype() error = %v", err)
|
||||
}
|
||||
if got != "mxfp8" {
|
||||
t.Fatalf("GetSafetensorsDtype() = %q, want mxfp8", got)
|
||||
if got != "nvfp4" {
|
||||
t.Fatalf("GetSafetensorsDtype() = %q, want nvfp4", got)
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user