import math
from typing import Callable, Optional, Union
import mlx.core as mx
from mlx.nn.layers.base import Module
from mlx.utils import tree_map_with_path
def _defaults_for_mode(mode, group_size, bits):
mode_defaults = {
"affine": (64, 4),
"mxfp4": (32, 4),
"nvfp4": (16, 4),
"mxfp8": (32, 8),
}
default_group_size, default_bits = mode_defaults[mode]
return group_size or default_group_size, bits or default_bits
def quantize(
model: Module,
group_size: int = None,
bits: int = None,
*,
mode: str = "affine",
quantize_input: bool = False,
class_predicate: Optional[Callable[[str, Module], Union[bool, dict]]] = None,
):
class_predicate = class_predicate or (lambda _, m: hasattr(m, "to_quantized"))
def _maybe_quantize(path, m):
if bool_or_params := class_predicate(path, m):
if hasattr(m, "to_quantized"):
if isinstance(bool_or_params, bool):
kwargs = {"group_size": group_size, "bits": bits, "mode": mode}
if quantize_input:
kwargs["quantize_input"] = quantize_input
return m.to_quantized(**kwargs)
elif isinstance(bool_or_params, dict):
if ("quantize_input" in bool_or_params) and not bool_or_params[
"quantize_input"
]:
bool_or_params.pop("quantize_input")
return m.to_quantized(**bool_or_params)
else:
raise ValueError(
"``class_predicate`` must return a bool"
" or a dict of parameters to pass to ``to_quantized``"
)
else:
raise ValueError(f"Unable to quantize model of type {type(m)}")
else:
return m
leaves = model.leaf_modules()
leaves = tree_map_with_path(_maybe_quantize, leaves, is_leaf=Module.is_module)
model.update_modules(leaves)
class QuantizedEmbedding(Module):
def __init__(
self,
num_embeddings: int,
dims: int,
group_size: int = None,
bits: int = None,
mode: str = "affine",
):
super().__init__()
self.group_size, self.bits = _defaults_for_mode(mode, group_size, bits)
self.mode = mode
scale = math.sqrt(1 / dims)
weight = mx.random.normal(shape=(num_embeddings, dims), scale=scale)
self.weight, self.scales, *biases = mx.quantize(
weight, group_size, bits, mode=mode
)
self.biases = biases[0] if biases else None
self.num_embeddings = num_embeddings
self.dims = dims
self.freeze()
def __call__(self, x):
biases = self.get("biases")
return mx.dequantize(
self["weight"][x],
scales=self["scales"][x],
biases=biases[x] if biases is not None else None,
group_size=self.group_size,
bits=self.bits,
mode=self.mode,
)
def as_linear(self, x):
return mx.quantized_matmul(
x,
self["weight"],
scales=self["scales"],
biases=self.get("biases"),
transpose=True,
group_size=self.group_size,
bits=self.bits,
mode=self.mode,
)
def _extra_repr(self):
return (
f"{self.num_embeddings}, {self.dims}, "
f"group_size={self.group_size}, bits={self.bits}, mode={self.mode}"
)
@classmethod
def from_embedding(
cls,
embedding_layer: Module,
group_size: int = None,
bits: int = None,
mode: str = "affine",
):
embedding_dims, dims = embedding_layer.weight.shape
ql = cls(embedding_dims, dims, group_size, bits, mode=mode)
ql.weight, ql.scales, *biases = mx.quantize(
embedding_layer.weight,
group_size,
bits,
mode=mode,
)
ql.biases = biases[0] if biases else None
return ql
class QuantizedLinear(Module):
def __init__(
self,
input_dims: int,
output_dims: int,
bias: bool = True,
group_size: int = None,
bits: int = None,
mode: str = "affine",
):
super().__init__()
self.group_size, self.bits = _defaults_for_mode(mode, group_size, bits)
self.mode = mode
scale = math.sqrt(1 / input_dims)
weight = mx.random.uniform(
low=-scale,
high=scale,
shape=(output_dims, input_dims),
)
self.weight, self.scales, *biases = mx.quantize(
weight, group_size, bits, mode=mode
)
self.biases = biases[0] if biases else None
if bias:
self.bias = mx.zeros((output_dims,))
self.freeze()
def _extra_repr(self):
out_dims, in_dims = self.weight.shape
in_dims = (in_dims * 32) // self.bits
return (
f"input_dims={in_dims}, output_dims={out_dims}, bias={'bias' in self}, "
f"group_size={self.group_size}, bits={self.bits}, mode={self.mode}"
)
def __call__(self, x):
x = mx.quantized_matmul(
x,
self["weight"],
scales=self["scales"],
biases=self.get("biases"),
transpose=True,
group_size=self.group_size,
bits=self.bits,
mode=self.mode,
)
if "bias" in self:
x = x + self["bias"]
return x
@classmethod
def from_linear(
cls,
linear_layer: Module,
group_size: int = None,
bits: int = None,
mode: str = "affine",
):
output_dims, input_dims = linear_layer.weight.shape
ql = cls(input_dims, output_dims, False, group_size, bits, mode=mode)
ql.weight, ql.scales, *biases = mx.quantize(
linear_layer.weight,
group_size,
bits,
mode=mode,
)
ql.biases = biases[0] if biases else None
if "bias" in linear_layer:
ql.bias = linear_layer.bias
return ql
class QQLinear(Module):
def __init__(
self,
input_dims: int,
output_dims: int,
group_size: int = None,
bits: int = None,
mode: str = "nvfp4",
):
super().__init__()
self.group_size, self.bits = _defaults_for_mode(mode, group_size, bits)
self.mode = mode
scale = math.sqrt(1 / input_dims)
self.weight = mx.random.uniform(
low=-scale,
high=scale,
shape=(output_dims, input_dims),
)
self._quantized = False
def _extra_repr(self):
out_dims, in_dims = self.weight.shape
if self.weight.dtype == mx.uint32:
in_dims = (in_dims * 32) // self.bits
return (
f"input_dims={in_dims}, output_dims={out_dims}, "
f"group_size={self.group_size}, bits={self.bits}, mode={self.mode}"
)
def quantize(self):
if not self._quantized:
self.weight, self.scales = mx.quantize(
self.weight,
self.group_size,
self.bits,
mode=self.mode,
)
self._quantized = True
def dequantize(self):
if self._quantized:
self.weight = mx.dequantize(
self.weight,
scales=self.scales,
group_size=self.group_size,
bits=self.bits,
mode=self.mode,
)
self.__delattr__("scales")
self._quantized = False
def _set_training_mode(self, mode: bool):
super()._set_training_mode(mode)
if self._training:
self.dequantize()
else:
self.quantize()
def __call__(self, x):
x = mx.qqmm(
x,
self["weight"],
scales=self.get("scales"),
group_size=self.group_size,
bits=self.bits,
mode=self.mode,
)
return x
@classmethod
def from_linear(
cls,
linear_layer: Module,
group_size: int = None,
bits: int = None,
mode: str = "nvfp4",
):
output_dims, input_dims = linear_layer.weight.shape if linear_layer.get("bias") is not None:
raise NotImplementedError("QQLinear does not support bias yet.")
ql = cls(input_dims, output_dims, group_size, bits, mode=mode)
ql.weight = linear_layer.weight
ql.train(linear_layer.training)
return ql