mirror of
https://github.com/ollama/ollama.git
synced 2026-07-23 09:10:53 -05:00
qwen3_5: apply nvfp4 global scales to the experts
Checkpoints quantized with NVIDIA's Model Optimizer carry a per-tensor scale on top of the group scales; MLX-quantized nvfp4 does not. The expert loaders dropped it, mis-scaling every expert output for those models. Apply it to the gather outputs on the quantized path and fold it into the weights when dequantizing.
This commit is contained in:
@@ -6,6 +6,7 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"math"
|
"math"
|
||||||
|
"slices"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
"github.com/ollama/ollama/x/mlxrunner/batch"
|
"github.com/ollama/ollama/x/mlxrunner/batch"
|
||||||
@@ -169,12 +170,20 @@ type SwitchMLP struct {
|
|||||||
DownGroupSize int
|
DownGroupSize int
|
||||||
GateUpMode string
|
GateUpMode string
|
||||||
DownMode string
|
DownMode string
|
||||||
|
|
||||||
|
// Tensor-level nvfp4 scales, [experts, 1, 1], applied to gather outputs.
|
||||||
|
GateUpGlobalScale, DownGlobalScale *mlx.Array
|
||||||
}
|
}
|
||||||
|
|
||||||
type stackedExpertWeights struct {
|
type stackedExpertWeights struct {
|
||||||
Weight *mlx.Array
|
Weight *mlx.Array
|
||||||
Scales *mlx.Array
|
Scales *mlx.Array
|
||||||
Biases *mlx.Array
|
Biases *mlx.Array
|
||||||
|
|
||||||
|
// Tensor-level nvfp4 scales as [experts, 1, 1], kept only for GatherQMM;
|
||||||
|
// dequantizing paths fold them into the weights.
|
||||||
|
GlobalScales *mlx.Array
|
||||||
|
|
||||||
Bits int
|
Bits int
|
||||||
GroupSize int
|
GroupSize int
|
||||||
Mode string
|
Mode string
|
||||||
@@ -470,6 +479,56 @@ func fuseExpertStacks(a, b *mlx.Array, axis int) *mlx.Array {
|
|||||||
return out
|
return out
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// tensorGlobalScale returns key's tensor-level nvfp4 scale; the import
|
||||||
|
// pipeline stores it as ModelOpt's weight_global_scale reciprocal.
|
||||||
|
func tensorGlobalScale(tensors map[string]*mlx.Array, key string) *mlx.Array {
|
||||||
|
if gs := tensors[key+".global_scale"]; gs != nil {
|
||||||
|
return gs
|
||||||
|
}
|
||||||
|
return tensors[key+".weight.global_scale"]
|
||||||
|
}
|
||||||
|
|
||||||
|
// expandExpertGlobalScale shapes scales as [experts, 1, 1] so Take with
|
||||||
|
// expert indices broadcasts against gather output rows.
|
||||||
|
func expandExpertGlobalScale(gs *mlx.Array, numExperts int) *mlx.Array {
|
||||||
|
if gs == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if gs.Size() == 1 {
|
||||||
|
gs = mlx.Add(mlx.Zeros(mlx.DTypeFloat32, numExperts, 1, 1), mlx.Reshape(gs, 1, 1, 1))
|
||||||
|
} else {
|
||||||
|
gs = mlx.Reshape(gs, int32(numExperts), 1, 1)
|
||||||
|
}
|
||||||
|
out := gs.Clone()
|
||||||
|
mlx.Eval(out)
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
func foldExpertGlobalScale(w, gs *mlx.Array) *mlx.Array {
|
||||||
|
if gs == nil {
|
||||||
|
return w
|
||||||
|
}
|
||||||
|
return mlx.Mul(w, gs).AsType(w.DType())
|
||||||
|
}
|
||||||
|
|
||||||
|
// sameExpertGlobalScales reports whether two per-expert scale stacks carry
|
||||||
|
// identical values (nil means absent).
|
||||||
|
func sameExpertGlobalScales(a, b *mlx.Array) bool {
|
||||||
|
if a == nil || b == nil {
|
||||||
|
return a == b
|
||||||
|
}
|
||||||
|
if a.Size() != b.Size() {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
av, bv := a.Floats(), b.Floats()
|
||||||
|
for i := range av {
|
||||||
|
if av[i] != bv[i] {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
// fuseGateUpProjections joins gate and up stacks along the output dimension,
|
// fuseGateUpProjections joins gate and up stacks along the output dimension,
|
||||||
// which is exact for quantized stacks: groups run along the input dimension.
|
// which is exact for quantized stacks: groups run along the input dimension.
|
||||||
func fuseGateUpProjections(gate, up *stackedExpertWeights) *stackedExpertWeights {
|
func fuseGateUpProjections(gate, up *stackedExpertWeights) *stackedExpertWeights {
|
||||||
@@ -481,13 +540,15 @@ func fuseGateUpProjections(gate, up *stackedExpertWeights) *stackedExpertWeights
|
|||||||
}
|
}
|
||||||
if gate.Scales != nil && up.Scales != nil &&
|
if gate.Scales != nil && up.Scales != nil &&
|
||||||
gate.Bits == up.Bits && gate.GroupSize == up.GroupSize && gate.Mode == up.Mode &&
|
gate.Bits == up.Bits && gate.GroupSize == up.GroupSize && gate.Mode == up.Mode &&
|
||||||
(gate.Biases == nil) == (up.Biases == nil) {
|
(gate.Biases == nil) == (up.Biases == nil) &&
|
||||||
|
sameExpertGlobalScales(gate.GlobalScales, up.GlobalScales) {
|
||||||
fused := &stackedExpertWeights{
|
fused := &stackedExpertWeights{
|
||||||
Weight: fuseExpertStacks(gate.Weight, up.Weight, 1),
|
Weight: fuseExpertStacks(gate.Weight, up.Weight, 1),
|
||||||
Scales: fuseExpertStacks(gate.Scales, up.Scales, 1),
|
Scales: fuseExpertStacks(gate.Scales, up.Scales, 1),
|
||||||
Bits: gate.Bits,
|
GlobalScales: gate.GlobalScales,
|
||||||
GroupSize: gate.GroupSize,
|
Bits: gate.Bits,
|
||||||
Mode: gate.Mode,
|
GroupSize: gate.GroupSize,
|
||||||
|
Mode: gate.Mode,
|
||||||
}
|
}
|
||||||
if gate.Biases != nil {
|
if gate.Biases != nil {
|
||||||
fused.Biases = fuseExpertStacks(gate.Biases, up.Biases, 1)
|
fused.Biases = fuseExpertStacks(gate.Biases, up.Biases, 1)
|
||||||
@@ -503,10 +564,12 @@ func fuseGateUpProjections(gate, up *stackedExpertWeights) *stackedExpertWeights
|
|||||||
if gate.Scales != nil {
|
if gate.Scales != nil {
|
||||||
gateWeight = mlx.Dequantize(gate.Weight, gate.Scales, gate.Biases, gate.GroupSize, gate.Bits, gate.Mode)
|
gateWeight = mlx.Dequantize(gate.Weight, gate.Scales, gate.Biases, gate.GroupSize, gate.Bits, gate.Mode)
|
||||||
}
|
}
|
||||||
|
gateWeight = foldExpertGlobalScale(gateWeight, gate.GlobalScales)
|
||||||
upWeight := up.Weight
|
upWeight := up.Weight
|
||||||
if up.Scales != nil {
|
if up.Scales != nil {
|
||||||
upWeight = mlx.Dequantize(up.Weight, up.Scales, up.Biases, up.GroupSize, up.Bits, up.Mode)
|
upWeight = mlx.Dequantize(up.Weight, up.Scales, up.Biases, up.GroupSize, up.Bits, up.Mode)
|
||||||
}
|
}
|
||||||
|
upWeight = foldExpertGlobalScale(upWeight, up.GlobalScales)
|
||||||
return &stackedExpertWeights{
|
return &stackedExpertWeights{
|
||||||
Weight: mlx.Concatenate([]*mlx.Array{gateWeight, upWeight}, 1),
|
Weight: mlx.Concatenate([]*mlx.Array{gateWeight, upWeight}, 1),
|
||||||
Bits: gate.Bits,
|
Bits: gate.Bits,
|
||||||
@@ -528,6 +591,7 @@ func loadStackedProjection(tensors map[string]*mlx.Array, cfg *Config, useQuanti
|
|||||||
}
|
}
|
||||||
|
|
||||||
qbiases := tensors[key+"_qbias"]
|
qbiases := tensors[key+"_qbias"]
|
||||||
|
globalScale := expandExpertGlobalScale(tensorGlobalScale(tensors, key), w.Dim(0))
|
||||||
groupSize, bits, mode := model.ResolveLinearQuantParams(
|
groupSize, bits, mode := model.ResolveLinearQuantParams(
|
||||||
cfg.QuantGroupSize,
|
cfg.QuantGroupSize,
|
||||||
cfg.QuantBits,
|
cfg.QuantBits,
|
||||||
@@ -539,12 +603,13 @@ func loadStackedProjection(tensors map[string]*mlx.Array, cfg *Config, useQuanti
|
|||||||
)
|
)
|
||||||
if useQuantized && supportsGatherQMM(mode, bits) {
|
if useQuantized && supportsGatherQMM(mode, bits) {
|
||||||
return &stackedExpertWeights{
|
return &stackedExpertWeights{
|
||||||
Weight: w,
|
Weight: w,
|
||||||
Scales: scales,
|
Scales: scales,
|
||||||
Biases: qbiases,
|
Biases: qbiases,
|
||||||
Bits: bits,
|
GlobalScales: globalScale,
|
||||||
GroupSize: groupSize,
|
Bits: bits,
|
||||||
Mode: mode,
|
GroupSize: groupSize,
|
||||||
|
Mode: mode,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -552,7 +617,7 @@ func loadStackedProjection(tensors map[string]*mlx.Array, cfg *Config, useQuanti
|
|||||||
slog.Warn("dequantizing expert weights: no gather kernel for format", "tensor", key, "mode", mode, "bits", bits)
|
slog.Warn("dequantizing expert weights: no gather kernel for format", "tensor", key, "mode", mode, "bits", bits)
|
||||||
}
|
}
|
||||||
return &stackedExpertWeights{
|
return &stackedExpertWeights{
|
||||||
Weight: mlx.Dequantize(w, scales, qbiases, groupSize, bits, mode),
|
Weight: foldExpertGlobalScale(mlx.Dequantize(w, scales, qbiases, groupSize, bits, mode), globalScale),
|
||||||
Bits: bits,
|
Bits: bits,
|
||||||
GroupSize: groupSize,
|
GroupSize: groupSize,
|
||||||
Mode: mode,
|
Mode: mode,
|
||||||
@@ -562,15 +627,17 @@ func loadStackedProjection(tensors map[string]*mlx.Array, cfg *Config, useQuanti
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// expertProjection is one expert's tensors for a projection; the quant fields
|
||||||
|
// are zero when the expert ships unquantized.
|
||||||
|
type expertProjection struct {
|
||||||
|
Weight, Scales, Biases, GlobalScale *mlx.Array
|
||||||
|
Bits, GroupSize int
|
||||||
|
Mode string
|
||||||
|
}
|
||||||
|
|
||||||
func collectPerExpertProjection(tensors map[string]*mlx.Array, cfg *Config, useQuantized bool, layerPrefix, proj string, numExperts int32) *stackedExpertWeights {
|
func collectPerExpertProjection(tensors map[string]*mlx.Array, cfg *Config, useQuantized bool, layerPrefix, proj string, numExperts int32) *stackedExpertWeights {
|
||||||
weights := make([]*mlx.Array, 0, numExperts)
|
experts := make([]expertProjection, 0, numExperts)
|
||||||
scales := make([]*mlx.Array, 0, numExperts)
|
consumedKeys := make([]string, 0, numExperts*4)
|
||||||
biases := make([]*mlx.Array, 0, numExperts)
|
|
||||||
consumedKeys := make([]string, 0, numExperts*3)
|
|
||||||
numDequantized := 0
|
|
||||||
bits := 0
|
|
||||||
groupSize := 0
|
|
||||||
mode := cfg.QuantMode
|
|
||||||
|
|
||||||
for e := range numExperts {
|
for e := range numExperts {
|
||||||
base := fmt.Sprintf("%s.mlp.experts.%d.%s", layerPrefix, e, proj)
|
base := fmt.Sprintf("%s.mlp.experts.%d.%s", layerPrefix, e, proj)
|
||||||
@@ -580,56 +647,90 @@ func collectPerExpertProjection(tensors map[string]*mlx.Array, cfg *Config, useQ
|
|||||||
}
|
}
|
||||||
consumedKeys = append(consumedKeys, key)
|
consumedKeys = append(consumedKeys, key)
|
||||||
|
|
||||||
s := tensors[key+"_scale"]
|
ex := expertProjection{Weight: w}
|
||||||
if s == nil {
|
if s := tensors[key+"_scale"]; s != nil {
|
||||||
weights = append(weights, w)
|
consumedKeys = append(consumedKeys, key+"_scale")
|
||||||
continue
|
ex.Scales = s
|
||||||
}
|
if ex.Biases = tensors[key+"_qbias"]; ex.Biases != nil {
|
||||||
consumedKeys = append(consumedKeys, key+"_scale")
|
consumedKeys = append(consumedKeys, key+"_qbias")
|
||||||
qb := tensors[key+"_qbias"]
|
|
||||||
if qb != nil {
|
|
||||||
consumedKeys = append(consumedKeys, key+"_qbias")
|
|
||||||
}
|
|
||||||
gs, b, m := model.ResolveLinearQuantParams(
|
|
||||||
cfg.QuantGroupSize,
|
|
||||||
cfg.QuantBits,
|
|
||||||
cfg.QuantMode,
|
|
||||||
cfg.TensorQuant,
|
|
||||||
key,
|
|
||||||
w,
|
|
||||||
s,
|
|
||||||
)
|
|
||||||
if bits == 0 {
|
|
||||||
bits = b
|
|
||||||
groupSize = gs
|
|
||||||
mode = m
|
|
||||||
}
|
|
||||||
if useQuantized && supportsGatherQMM(m, b) {
|
|
||||||
weights = append(weights, w)
|
|
||||||
scales = append(scales, s)
|
|
||||||
if qb != nil {
|
|
||||||
biases = append(biases, qb)
|
|
||||||
}
|
}
|
||||||
} else {
|
if ex.GlobalScale = tensorGlobalScale(tensors, key); ex.GlobalScale != nil {
|
||||||
weights = append(weights, mlx.Dequantize(w, s, qb, gs, b, m))
|
consumedKeys = append(consumedKeys, key+".global_scale", key+".weight.global_scale")
|
||||||
numDequantized++
|
}
|
||||||
|
ex.GroupSize, ex.Bits, ex.Mode = model.ResolveLinearQuantParams(
|
||||||
|
cfg.QuantGroupSize,
|
||||||
|
cfg.QuantBits,
|
||||||
|
cfg.QuantMode,
|
||||||
|
cfg.TensorQuant,
|
||||||
|
key,
|
||||||
|
w,
|
||||||
|
s,
|
||||||
|
)
|
||||||
}
|
}
|
||||||
|
experts = append(experts, ex)
|
||||||
}
|
}
|
||||||
if useQuantized && numDequantized > 0 {
|
if len(experts) == 0 {
|
||||||
slog.Warn("dequantizing expert weights: no gather kernel for format",
|
|
||||||
"tensor", fmt.Sprintf("%s.mlp.experts.*.%s", layerPrefix, proj), "mode", mode, "bits", bits)
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(weights) == 0 {
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
out := &stackedExpertWeights{Weight: stackAndClone(weights), Bits: bits, GroupSize: groupSize, Mode: mode}
|
// The gather path needs every expert quantized in one shared format;
|
||||||
if len(scales) == len(weights) {
|
// otherwise dequantize them all so the stack is uniformly fp.
|
||||||
out.Scales = stackAndClone(scales)
|
first := experts[0]
|
||||||
|
quantized := useQuantized
|
||||||
|
for _, ex := range experts {
|
||||||
|
if ex.Scales == nil || !supportsGatherQMM(ex.Mode, ex.Bits) ||
|
||||||
|
ex.Bits != first.Bits || ex.GroupSize != first.GroupSize || ex.Mode != first.Mode ||
|
||||||
|
(ex.Biases == nil) != (first.Biases == nil) {
|
||||||
|
quantized = false
|
||||||
|
break
|
||||||
|
}
|
||||||
}
|
}
|
||||||
if len(biases) == len(weights) {
|
|
||||||
out.Biases = stackAndClone(biases)
|
parts := make([]*mlx.Array, len(experts))
|
||||||
|
out := &stackedExpertWeights{Mode: cfg.QuantMode}
|
||||||
|
if quantized {
|
||||||
|
out.Bits, out.GroupSize, out.Mode = first.Bits, first.GroupSize, first.Mode
|
||||||
|
for i, ex := range experts {
|
||||||
|
parts[i] = ex.Weight
|
||||||
|
}
|
||||||
|
out.Weight = stackAndClone(parts)
|
||||||
|
for i, ex := range experts {
|
||||||
|
parts[i] = ex.Scales
|
||||||
|
}
|
||||||
|
out.Scales = stackAndClone(parts)
|
||||||
|
if first.Biases != nil {
|
||||||
|
for i, ex := range experts {
|
||||||
|
parts[i] = ex.Biases
|
||||||
|
}
|
||||||
|
out.Biases = stackAndClone(parts)
|
||||||
|
}
|
||||||
|
if slices.ContainsFunc(experts, func(ex expertProjection) bool { return ex.GlobalScale != nil }) {
|
||||||
|
for i, ex := range experts {
|
||||||
|
if ex.GlobalScale == nil {
|
||||||
|
// An expert without a tensor-level scale is unscaled.
|
||||||
|
parts[i] = mlx.FromValues([]float32{1}, 1)
|
||||||
|
} else {
|
||||||
|
parts[i] = mlx.Reshape(ex.GlobalScale.AsType(mlx.DTypeFloat32), 1)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
out.GlobalScales = expandExpertGlobalScale(stackAndClone(parts), len(experts))
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
if useQuantized && slices.ContainsFunc(experts, func(ex expertProjection) bool { return ex.Scales != nil }) {
|
||||||
|
slog.Warn("dequantizing expert weights: unsupported or mixed formats",
|
||||||
|
"tensor", fmt.Sprintf("%s.mlp.experts.*.%s", layerPrefix, proj))
|
||||||
|
}
|
||||||
|
for i, ex := range experts {
|
||||||
|
if ex.Scales == nil {
|
||||||
|
parts[i] = ex.Weight
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if out.Bits == 0 {
|
||||||
|
out.Bits, out.GroupSize, out.Mode = ex.Bits, ex.GroupSize, ex.Mode
|
||||||
|
}
|
||||||
|
parts[i] = foldExpertGlobalScale(mlx.Dequantize(ex.Weight, ex.Scales, ex.Biases, ex.GroupSize, ex.Bits, ex.Mode), ex.GlobalScale)
|
||||||
|
}
|
||||||
|
out.Weight = stackAndClone(parts)
|
||||||
}
|
}
|
||||||
freeTensorKeys(tensors, consumedKeys...)
|
freeTensorKeys(tensors, consumedKeys...)
|
||||||
return out
|
return out
|
||||||
@@ -653,6 +754,7 @@ func combinedGateUpProjection(tensors map[string]*mlx.Array, cfg *Config, useQua
|
|||||||
}
|
}
|
||||||
|
|
||||||
qbiases := tensors[key+"_qbias"]
|
qbiases := tensors[key+"_qbias"]
|
||||||
|
globalScale := expandExpertGlobalScale(tensorGlobalScale(tensors, key), gateUp.Dim(0))
|
||||||
groupSize, bits, mode := model.ResolveLinearQuantParams(
|
groupSize, bits, mode := model.ResolveLinearQuantParams(
|
||||||
cfg.QuantGroupSize,
|
cfg.QuantGroupSize,
|
||||||
cfg.QuantBits,
|
cfg.QuantBits,
|
||||||
@@ -664,19 +766,20 @@ func combinedGateUpProjection(tensors map[string]*mlx.Array, cfg *Config, useQua
|
|||||||
)
|
)
|
||||||
if useQuantized && supportsGatherQMM(mode, bits) {
|
if useQuantized && supportsGatherQMM(mode, bits) {
|
||||||
return &stackedExpertWeights{
|
return &stackedExpertWeights{
|
||||||
Weight: gateUp,
|
Weight: gateUp,
|
||||||
Scales: scales,
|
Scales: scales,
|
||||||
Biases: qbiases,
|
Biases: qbiases,
|
||||||
Bits: bits,
|
GlobalScales: globalScale,
|
||||||
GroupSize: groupSize,
|
Bits: bits,
|
||||||
Mode: mode,
|
GroupSize: groupSize,
|
||||||
|
Mode: mode,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if useQuantized {
|
if useQuantized {
|
||||||
slog.Warn("dequantizing expert weights: no gather kernel for format", "tensor", key, "mode", mode, "bits", bits)
|
slog.Warn("dequantizing expert weights: no gather kernel for format", "tensor", key, "mode", mode, "bits", bits)
|
||||||
}
|
}
|
||||||
return &stackedExpertWeights{
|
return &stackedExpertWeights{
|
||||||
Weight: mlx.Dequantize(gateUp, scales, qbiases, groupSize, bits, mode),
|
Weight: foldExpertGlobalScale(mlx.Dequantize(gateUp, scales, qbiases, groupSize, bits, mode), globalScale),
|
||||||
Bits: bits,
|
Bits: bits,
|
||||||
GroupSize: groupSize,
|
GroupSize: groupSize,
|
||||||
Mode: mode,
|
Mode: mode,
|
||||||
@@ -740,6 +843,7 @@ func loadSwitchMLP(tensors map[string]*mlx.Array, cfg *Config, useQuantized bool
|
|||||||
switchMLP.GateUpBits = gateUpW.Bits
|
switchMLP.GateUpBits = gateUpW.Bits
|
||||||
switchMLP.GateUpGroupSize = gateUpW.GroupSize
|
switchMLP.GateUpGroupSize = gateUpW.GroupSize
|
||||||
switchMLP.GateUpMode = gateUpW.Mode
|
switchMLP.GateUpMode = gateUpW.Mode
|
||||||
|
switchMLP.GateUpGlobalScale = gateUpW.GlobalScales
|
||||||
} else {
|
} else {
|
||||||
switchMLP.GateUpWeight = transposeExpertWeightForGatherMM(gateUpW.Weight)
|
switchMLP.GateUpWeight = transposeExpertWeightForGatherMM(gateUpW.Weight)
|
||||||
}
|
}
|
||||||
@@ -750,6 +854,7 @@ func loadSwitchMLP(tensors map[string]*mlx.Array, cfg *Config, useQuantized bool
|
|||||||
switchMLP.DownBits = downW.Bits
|
switchMLP.DownBits = downW.Bits
|
||||||
switchMLP.DownGroupSize = downW.GroupSize
|
switchMLP.DownGroupSize = downW.GroupSize
|
||||||
switchMLP.DownMode = downW.Mode
|
switchMLP.DownMode = downW.Mode
|
||||||
|
switchMLP.DownGlobalScale = downW.GlobalScales
|
||||||
} else {
|
} else {
|
||||||
switchMLP.DownWeight = transposeExpertWeightForGatherMM(downW.Weight)
|
switchMLP.DownWeight = transposeExpertWeightForGatherMM(downW.Weight)
|
||||||
}
|
}
|
||||||
@@ -1154,6 +1259,9 @@ func (s *SwitchMLP) Forward(x *mlx.Array, indices *mlx.Array, cfg *Config) *mlx.
|
|||||||
} else {
|
} else {
|
||||||
gateUp = mlx.GatherMM(xFlat, s.GateUpWeight, nil, idxFlat, doSort)
|
gateUp = mlx.GatherMM(xFlat, s.GateUpWeight, nil, idxFlat, doSort)
|
||||||
}
|
}
|
||||||
|
if s.GateUpGlobalScale != nil {
|
||||||
|
gateUp = mlx.Mul(gateUp, mlx.Take(s.GateUpGlobalScale, idxFlat, 0)).AsType(xFlat.DType())
|
||||||
|
}
|
||||||
gate, up := splitLastAxisHalves(gateUp)
|
gate, up := splitLastAxisHalves(gateUp)
|
||||||
hidden := mlx.SwiGLU(gate, up)
|
hidden := mlx.SwiGLU(gate, up)
|
||||||
if s.DownWeightQ != nil {
|
if s.DownWeightQ != nil {
|
||||||
@@ -1162,6 +1270,9 @@ func (s *SwitchMLP) Forward(x *mlx.Array, indices *mlx.Array, cfg *Config) *mlx.
|
|||||||
} else {
|
} else {
|
||||||
down = mlx.GatherMM(hidden, s.DownWeight, nil, idxFlat, doSort)
|
down = mlx.GatherMM(hidden, s.DownWeight, nil, idxFlat, doSort)
|
||||||
}
|
}
|
||||||
|
if s.DownGlobalScale != nil {
|
||||||
|
down = mlx.Mul(down, mlx.Take(s.DownGlobalScale, idxFlat, 0)).AsType(hidden.DType())
|
||||||
|
}
|
||||||
|
|
||||||
if doSort {
|
if doSort {
|
||||||
down = mlx.Reshape(mlx.Take(mlx.Squeeze(mlx.Squeeze(down, 2), 1), invOrder, 0), B*L, topK, cfg.HiddenSize)
|
down = mlx.Reshape(mlx.Take(mlx.Squeeze(mlx.Squeeze(down, 2), 1), invOrder, 0), B*L, topK, cfg.HiddenSize)
|
||||||
|
|||||||
296
x/models/qwen3_5/qwen3_5_moe_test.go
Normal file
296
x/models/qwen3_5/qwen3_5_moe_test.go
Normal file
@@ -0,0 +1,296 @@
|
|||||||
|
package qwen3_5
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"math"
|
||||||
|
"runtime"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/ollama/ollama/x/mlxrunner/mlx"
|
||||||
|
)
|
||||||
|
|
||||||
|
// MLX streams are bound to the thread that created them.
|
||||||
|
func useMLXTestThread(t *testing.T) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
runtime.LockOSThread()
|
||||||
|
initialized := false
|
||||||
|
t.Cleanup(func() {
|
||||||
|
if initialized {
|
||||||
|
mlx.Sweep()
|
||||||
|
mlx.ClearCache()
|
||||||
|
}
|
||||||
|
runtime.UnlockOSThread()
|
||||||
|
})
|
||||||
|
|
||||||
|
if err := mlx.CheckInit(); err != nil {
|
||||||
|
t.Skipf("MLX not available: %v", err)
|
||||||
|
}
|
||||||
|
initialized = true
|
||||||
|
if mlx.GPUIsAvailable() {
|
||||||
|
mlx.SetDefaultDeviceGPU()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func moeTestConfig() *Config {
|
||||||
|
return &Config{
|
||||||
|
NumExperts: 8,
|
||||||
|
NumExpertsPerTok: 2,
|
||||||
|
HiddenSize: 128,
|
||||||
|
QuantGroupSize: 32,
|
||||||
|
QuantBits: 8,
|
||||||
|
QuantMode: "affine",
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func patternValues(n int, seed float64) []float32 {
|
||||||
|
vals := make([]float32, n)
|
||||||
|
for i := range vals {
|
||||||
|
vals[i] = float32(math.Sin(seed + float64(i)*0.37))
|
||||||
|
}
|
||||||
|
return vals
|
||||||
|
}
|
||||||
|
|
||||||
|
func expertSlice(a *mlx.Array, e int32) *mlx.Array {
|
||||||
|
dims := a.Dims()
|
||||||
|
return mlx.Squeeze(mlx.SliceStartStop(a, []int32{e, 0, 0}, []int32{e + 1, int32(dims[1]), int32(dims[2])}), 0)
|
||||||
|
}
|
||||||
|
|
||||||
|
type moeTestWeights struct {
|
||||||
|
gateQ, gateS, gateB *mlx.Array
|
||||||
|
upQ, upS, upB *mlx.Array
|
||||||
|
downQ, downS, downB *mlx.Array
|
||||||
|
}
|
||||||
|
|
||||||
|
func makeMoETestWeights(cfg *Config) moeTestWeights {
|
||||||
|
E, I, H := int(cfg.NumExperts), 64, int(cfg.HiddenSize)
|
||||||
|
gate := mlx.FromValues(patternValues(E*I*H, 1), E, I, H)
|
||||||
|
up := mlx.FromValues(patternValues(E*I*H, 2), E, I, H)
|
||||||
|
down := mlx.FromValues(patternValues(E*H*I, 3), E, H, I)
|
||||||
|
var w moeTestWeights
|
||||||
|
gs, bits, mode := cfg.QuantGroupSize, cfg.QuantBits, cfg.QuantMode
|
||||||
|
w.gateQ, w.gateS, w.gateB = mlx.Quantize(gate, gs, bits, mode)
|
||||||
|
w.upQ, w.upS, w.upB = mlx.Quantize(up, gs, bits, mode)
|
||||||
|
w.downQ, w.downS, w.downB = mlx.Quantize(down, gs, bits, mode)
|
||||||
|
return w
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w moeTestWeights) dequantized(cfg *Config) (gate, up, down *mlx.Array) {
|
||||||
|
gs, bits, mode := cfg.QuantGroupSize, cfg.QuantBits, cfg.QuantMode
|
||||||
|
return mlx.Dequantize(w.gateQ, w.gateS, w.gateB, gs, bits, mode),
|
||||||
|
mlx.Dequantize(w.upQ, w.upS, w.upB, gs, bits, mode),
|
||||||
|
mlx.Dequantize(w.downQ, w.downS, w.downB, gs, bits, mode)
|
||||||
|
}
|
||||||
|
|
||||||
|
func moeForward(t *testing.T, tensors map[string]*mlx.Array, cfg *Config, L int32) (*SwitchMLP, []float32) {
|
||||||
|
t.Helper()
|
||||||
|
mlp, err := loadSwitchMLP(tensors, cfg, true, "model.layers.0")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("loadSwitchMLP: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
x := mlx.FromValues(patternValues(int(L*cfg.HiddenSize), 9), 1, int(L), int(cfg.HiddenSize))
|
||||||
|
idx := make([]uint32, L*cfg.NumExpertsPerTok)
|
||||||
|
for i := range idx {
|
||||||
|
idx[i] = uint32((i*7 + 3) % int(cfg.NumExperts))
|
||||||
|
}
|
||||||
|
indices := mlx.FromValues(idx, 1, int(L), int(cfg.NumExpertsPerTok))
|
||||||
|
|
||||||
|
out := mlp.Forward(x, indices, cfg)
|
||||||
|
mlx.Eval(out)
|
||||||
|
return mlp, out.Floats()
|
||||||
|
}
|
||||||
|
|
||||||
|
func maxAbsDiff(a, b []float32) float64 {
|
||||||
|
var m float64
|
||||||
|
for i := range a {
|
||||||
|
if d := math.Abs(float64(a[i]) - float64(b[i])); d > m {
|
||||||
|
m = d
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return m
|
||||||
|
}
|
||||||
|
|
||||||
|
// Every supported checkpoint layout must produce the same routed-expert
|
||||||
|
// output once normalized to the packed gate_up representation.
|
||||||
|
func TestSwitchMLPSourceLayoutEquivalence(t *testing.T) {
|
||||||
|
useMLXTestThread(t)
|
||||||
|
|
||||||
|
cfg := moeTestConfig()
|
||||||
|
w := makeMoETestWeights(cfg)
|
||||||
|
prefix := "model.layers.0"
|
||||||
|
|
||||||
|
packed := func() map[string]*mlx.Array {
|
||||||
|
return map[string]*mlx.Array{
|
||||||
|
prefix + ".mlp.experts.gate_up_proj.weight": mlx.Concatenate([]*mlx.Array{w.gateQ, w.upQ}, 1),
|
||||||
|
prefix + ".mlp.experts.gate_up_proj.weight_scale": mlx.Concatenate([]*mlx.Array{w.gateS, w.upS}, 1),
|
||||||
|
prefix + ".mlp.experts.gate_up_proj.weight_qbias": mlx.Concatenate([]*mlx.Array{w.gateB, w.upB}, 1),
|
||||||
|
prefix + ".mlp.experts.down_proj.weight": w.downQ,
|
||||||
|
prefix + ".mlp.experts.down_proj.weight_scale": w.downS,
|
||||||
|
prefix + ".mlp.experts.down_proj.weight_qbias": w.downB,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
separate := func() map[string]*mlx.Array {
|
||||||
|
return map[string]*mlx.Array{
|
||||||
|
prefix + ".mlp.switch_mlp.gate_proj.weight": w.gateQ,
|
||||||
|
prefix + ".mlp.switch_mlp.gate_proj.weight_scale": w.gateS,
|
||||||
|
prefix + ".mlp.switch_mlp.gate_proj.weight_qbias": w.gateB,
|
||||||
|
prefix + ".mlp.switch_mlp.up_proj.weight": w.upQ,
|
||||||
|
prefix + ".mlp.switch_mlp.up_proj.weight_scale": w.upS,
|
||||||
|
prefix + ".mlp.switch_mlp.up_proj.weight_qbias": w.upB,
|
||||||
|
prefix + ".mlp.switch_mlp.down_proj.weight": w.downQ,
|
||||||
|
prefix + ".mlp.switch_mlp.down_proj.weight_scale": w.downS,
|
||||||
|
prefix + ".mlp.switch_mlp.down_proj.weight_qbias": w.downB,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
perExpert := func() map[string]*mlx.Array {
|
||||||
|
tensors := map[string]*mlx.Array{}
|
||||||
|
for e := range cfg.NumExperts {
|
||||||
|
for proj, parts := range map[string][3]*mlx.Array{
|
||||||
|
"gate_proj": {w.gateQ, w.gateS, w.gateB},
|
||||||
|
"up_proj": {w.upQ, w.upS, w.upB},
|
||||||
|
"down_proj": {w.downQ, w.downS, w.downB},
|
||||||
|
} {
|
||||||
|
base := fmt.Sprintf("%s.mlp.experts.%d.%s.weight", prefix, e, proj)
|
||||||
|
tensors[base] = expertSlice(parts[0], e)
|
||||||
|
tensors[base+"_scale"] = expertSlice(parts[1], e)
|
||||||
|
tensors[base+"_qbias"] = expertSlice(parts[2], e)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return tensors
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, L := range []int32{2, 64} { // below and above the sorted-gather threshold
|
||||||
|
mlp, want := moeForward(t, packed(), cfg, L)
|
||||||
|
if mlp.GateUpWeightQ == nil || mlp.DownWeightQ == nil {
|
||||||
|
t.Fatalf("L=%d: packed layout did not take the quantized path", L)
|
||||||
|
}
|
||||||
|
for name, tensors := range map[string]func() map[string]*mlx.Array{"separate": separate, "per_expert": perExpert} {
|
||||||
|
mlp, got := moeForward(t, tensors(), cfg, L)
|
||||||
|
if mlp.GateUpWeightQ == nil || mlp.DownWeightQ == nil {
|
||||||
|
t.Fatalf("L=%d %s: layout did not take the quantized path", L, name)
|
||||||
|
}
|
||||||
|
if d := maxAbsDiff(want, got); d != 0 {
|
||||||
|
t.Errorf("L=%d %s: output differs from packed layout, max abs diff %g", L, name, d)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Tensor-level scales must be applied to gather outputs on the quantized
|
||||||
|
// path and folded into the weights on every fp path.
|
||||||
|
func TestSwitchMLPGlobalScales(t *testing.T) {
|
||||||
|
useMLXTestThread(t)
|
||||||
|
|
||||||
|
cfg := moeTestConfig()
|
||||||
|
w := makeMoETestWeights(cfg)
|
||||||
|
deqGate, deqUp, deqDown := w.dequantized(cfg)
|
||||||
|
prefix := "model.layers.0"
|
||||||
|
|
||||||
|
// Sorted gathers run on reduced-precision NAX kernels.
|
||||||
|
const tolerance = 0.1
|
||||||
|
|
||||||
|
const gsGate, gsUp, gsDown = 0.5, 0.25, 2.0
|
||||||
|
separate := func() map[string]*mlx.Array {
|
||||||
|
return map[string]*mlx.Array{
|
||||||
|
prefix + ".mlp.switch_mlp.gate_proj.weight": w.gateQ,
|
||||||
|
prefix + ".mlp.switch_mlp.gate_proj.weight_scale": w.gateS,
|
||||||
|
prefix + ".mlp.switch_mlp.gate_proj.weight_qbias": w.gateB,
|
||||||
|
prefix + ".mlp.switch_mlp.gate_proj.weight.global_scale": mlx.FromValues([]float32{gsGate}, 1),
|
||||||
|
prefix + ".mlp.switch_mlp.up_proj.weight": w.upQ,
|
||||||
|
prefix + ".mlp.switch_mlp.up_proj.weight_scale": w.upS,
|
||||||
|
prefix + ".mlp.switch_mlp.up_proj.weight_qbias": w.upB,
|
||||||
|
prefix + ".mlp.switch_mlp.up_proj.weight.global_scale": mlx.FromValues([]float32{gsUp}, 1),
|
||||||
|
prefix + ".mlp.switch_mlp.down_proj.weight": w.downQ,
|
||||||
|
prefix + ".mlp.switch_mlp.down_proj.weight_scale": w.downS,
|
||||||
|
prefix + ".mlp.switch_mlp.down_proj.weight_qbias": w.downB,
|
||||||
|
prefix + ".mlp.switch_mlp.down_proj.weight.global_scale": mlx.FromValues([]float32{gsDown}, 1),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Packed checkpoints carry one scale covering both gate_up halves.
|
||||||
|
packed := func() map[string]*mlx.Array {
|
||||||
|
return map[string]*mlx.Array{
|
||||||
|
prefix + ".mlp.experts.gate_up_proj.weight": mlx.Concatenate([]*mlx.Array{w.gateQ, w.upQ}, 1),
|
||||||
|
prefix + ".mlp.experts.gate_up_proj.weight_scale": mlx.Concatenate([]*mlx.Array{w.gateS, w.upS}, 1),
|
||||||
|
prefix + ".mlp.experts.gate_up_proj.weight_qbias": mlx.Concatenate([]*mlx.Array{w.gateB, w.upB}, 1),
|
||||||
|
prefix + ".mlp.experts.gate_up_proj.weight.global_scale": mlx.FromValues([]float32{gsGate}, 1),
|
||||||
|
prefix + ".mlp.experts.down_proj.weight": w.downQ,
|
||||||
|
prefix + ".mlp.experts.down_proj.weight_scale": w.downS,
|
||||||
|
prefix + ".mlp.experts.down_proj.weight_qbias": w.downB,
|
||||||
|
prefix + ".mlp.experts.down_proj.weight.global_scale": mlx.FromValues([]float32{gsDown}, 1),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// References: the same weights dequantized with the scales pre-folded.
|
||||||
|
foldedSeparate := func() map[string]*mlx.Array {
|
||||||
|
return map[string]*mlx.Array{
|
||||||
|
prefix + ".mlp.switch_mlp.gate_proj.weight": mlx.MulScalar(deqGate, gsGate),
|
||||||
|
prefix + ".mlp.switch_mlp.up_proj.weight": mlx.MulScalar(deqUp, gsUp),
|
||||||
|
prefix + ".mlp.switch_mlp.down_proj.weight": mlx.MulScalar(deqDown, gsDown),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
foldedPacked := func() map[string]*mlx.Array {
|
||||||
|
return map[string]*mlx.Array{
|
||||||
|
prefix + ".mlp.switch_mlp.gate_proj.weight": mlx.MulScalar(deqGate, gsGate),
|
||||||
|
prefix + ".mlp.switch_mlp.up_proj.weight": mlx.MulScalar(deqUp, gsGate),
|
||||||
|
prefix + ".mlp.switch_mlp.down_proj.weight": mlx.MulScalar(deqDown, gsDown),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Quantized gate/up with fp down: gate_up stays on the gather kernel
|
||||||
|
// with its scales; only down runs fp.
|
||||||
|
mixed := func() map[string]*mlx.Array {
|
||||||
|
tensors := separate()
|
||||||
|
delete(tensors, prefix+".mlp.switch_mlp.gate_proj.weight.global_scale")
|
||||||
|
delete(tensors, prefix+".mlp.switch_mlp.up_proj.weight.global_scale")
|
||||||
|
delete(tensors, prefix+".mlp.switch_mlp.down_proj.weight_scale")
|
||||||
|
delete(tensors, prefix+".mlp.switch_mlp.down_proj.weight_qbias")
|
||||||
|
delete(tensors, prefix+".mlp.switch_mlp.down_proj.weight.global_scale")
|
||||||
|
tensors[prefix+".mlp.switch_mlp.down_proj.weight"] = mlx.MulScalar(deqDown, gsDown)
|
||||||
|
return tensors
|
||||||
|
}
|
||||||
|
foldedMixed := func() map[string]*mlx.Array {
|
||||||
|
return map[string]*mlx.Array{
|
||||||
|
prefix + ".mlp.switch_mlp.gate_proj.weight": deqGate,
|
||||||
|
prefix + ".mlp.switch_mlp.up_proj.weight": deqUp,
|
||||||
|
prefix + ".mlp.switch_mlp.down_proj.weight": mlx.MulScalar(deqDown, gsDown),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, L := range []int32{2, 64} {
|
||||||
|
mlpF, want := moeForward(t, foldedSeparate(), cfg, L)
|
||||||
|
if mlpF.GateUpWeightQ != nil || mlpF.GateUpGlobalScale != nil || mlpF.DownGlobalScale != nil {
|
||||||
|
t.Fatalf("L=%d: fp layout must fold tensor-level scales into the weights", L)
|
||||||
|
}
|
||||||
|
|
||||||
|
mlpQ, got := moeForward(t, separate(), cfg, L)
|
||||||
|
if mlpQ.GateUpWeightQ != nil || mlpQ.GateUpGlobalScale != nil {
|
||||||
|
t.Fatalf("L=%d: differing gate/up scales must decline quantized fusion", L)
|
||||||
|
}
|
||||||
|
if mlpQ.DownWeightQ == nil || mlpQ.DownGlobalScale == nil {
|
||||||
|
t.Fatalf("L=%d: down must keep its tensor-level scale for the gather path", L)
|
||||||
|
}
|
||||||
|
if d := maxAbsDiff(want, got); d > tolerance {
|
||||||
|
t.Errorf("L=%d: quantized output diverges from folded fp reference, max abs diff %g", L, d)
|
||||||
|
}
|
||||||
|
|
||||||
|
mlpM, gotMixed := moeForward(t, mixed(), cfg, L)
|
||||||
|
if mlpM.GateUpWeightQ == nil || mlpM.GateUpGlobalScale != nil || mlpM.DownWeightQ != nil || mlpM.DownGlobalScale != nil {
|
||||||
|
t.Fatalf("L=%d: mixed layout must keep gate_up quantized and run down fp", L)
|
||||||
|
}
|
||||||
|
_, wantMixed := moeForward(t, foldedMixed(), cfg, L)
|
||||||
|
if d := maxAbsDiff(wantMixed, gotMixed); d > tolerance {
|
||||||
|
t.Errorf("L=%d: mixed layout diverges from fp reference, max abs diff %g", L, d)
|
||||||
|
}
|
||||||
|
|
||||||
|
mlpP, gotPacked := moeForward(t, packed(), cfg, L)
|
||||||
|
if mlpP.GateUpWeightQ == nil || mlpP.GateUpGlobalScale == nil || mlpP.DownGlobalScale == nil {
|
||||||
|
t.Fatalf("L=%d: packed layout must keep tensor-level scales for the gather path", L)
|
||||||
|
}
|
||||||
|
_, wantPacked := moeForward(t, foldedPacked(), cfg, L)
|
||||||
|
if d := maxAbsDiff(wantPacked, gotPacked); d > tolerance {
|
||||||
|
t.Errorf("L=%d: packed output diverges from folded fp reference, max abs diff %g", L, d)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user