QQ linear (#2931)
This commit is contained in:
@@ -87,7 +87,12 @@ from mlx.nn.layers.pooling import (
|
||||
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,
|
||||
|
||||
@@ -559,6 +559,9 @@ class Module(dict):
|
||||
_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 @@ class Module(dict):
|
||||
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:
|
||||
|
||||
@@ -277,3 +277,128 @@ class QuantizedLinear(Module):
|
||||
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
|
||||
|
||||
Reference in New Issue
Block a user