diff --git a/python/mlx/nn/layers/__init__.py b/python/mlx/nn/layers/__init__.py index ea2d3029d8..c2fba58347 100644 --- a/python/mlx/nn/layers/__init__.py +++ b/python/mlx/nn/layers/__init__.py @@ -87,7 +87,12 @@ MaxPool3d, ) from mlx.nn.layers.positional_encoding import ALiBi, RoPE, SinusoidalPositionalEncoding -from mlx.nn.layers.quantized import QuantizedEmbedding, QuantizedLinear, quantize +from mlx.nn.layers.quantized import ( + QQLinear, + QuantizedEmbedding, + QuantizedLinear, + quantize, +) from mlx.nn.layers.recurrent import GRU, LSTM, RNN from mlx.nn.layers.transformer import ( MultiHeadAttention, diff --git a/python/mlx/nn/layers/base.py b/python/mlx/nn/layers/base.py index ce009d9f83..fca65d787d 100644 --- a/python/mlx/nn/layers/base.py +++ b/python/mlx/nn/layers/base.py @@ -559,6 +559,9 @@ def _unfreeze_impl(_, m): _unfreeze_impl("", self) return self + def _set_training_mode(self, mode: bool) -> None: + self._training = mode + def train(self, mode: bool = True) -> Module: """Set the model in or out of training mode. @@ -573,10 +576,8 @@ def train(self, mode: bool = True) -> Module: The module instance after updating the training mode. """ - def _set_train(_, m): - m._training = mode + self.apply_to_modules(lambda _, m: m._set_training_mode(mode)) - self.apply_to_modules(_set_train) return self def eval(self) -> Module: diff --git a/python/mlx/nn/layers/quantized.py b/python/mlx/nn/layers/quantized.py index f762847e4d..1c98706daf 100644 --- a/python/mlx/nn/layers/quantized.py +++ b/python/mlx/nn/layers/quantized.py @@ -277,3 +277,128 @@ def from_linear( ql.bias = linear_layer.bias return ql + + +class QQLinear(Module): + """Quantizes the input and applies an affine transformation using quantized weights. + + Two use cases are supported: + + 1) **Eval**: The weights are frozen and stored in quantized form together with + their scales (``self.weight`` is quantized and ``self.scales`` is provided). + 2) **Train**: The weights are stored in higher precision and are quantized on + the fly during computation so that gradients with respect to the weights + can be computed. + + To switch between the two cases, use ``layer.eval()`` and ``layer.train()`` respectively. + + Compared to the :class:`mlx.nn.QuantizedLinear` layer, this layer + quantizes the input as well and includes weights in gradient computations. + + :obj:`QQLinear` also provides: + - the class method :meth:`from_linear` to convert :class:`mlx.nn.Linear` + layers to :obj:`QQLinear` layers. + + Note: This layer does not support a bias term yet. + + Args: + input_dims (int): The dimensionality of the input features. + output_dims (int): The dimensionality of the output features. + group_size (Optional[int]): The group size to use for the quantized weight. + See :func:`~mlx.core.quantize`. Default: ``None``. + bits (Optional[int]): The bit width to use for the quantized weight. + See :func:`~mlx.core.quantize`. Default: ``None``. + mode (Optional[str]): The quantization method to use (see + :func:`mlx.core.quantize`). Currently, only ``"nvfp4"`` and ``"mxfp8"`` + are supported. Default: ``"nvfp4"``. + """ + + def __init__( + self, + input_dims: int, + output_dims: int, + group_size: int = None, + bits: int = None, + mode: str = "nvfp4", + ): + super().__init__() + + # Quantization config + 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 *= 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", + ): + """Create a :obj:`QQLinear` layer from a :obj:`Linear` layer.""" + output_dims, input_dims = linear_layer.weight.shape # (N,K) + 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