Fix tolerance in de-/quantization test

This commit is contained in:
Angelos Katharopoulos 2023-12-26 17:59:33 -08:00
parent fc4e5b476b
commit f0bf2bf09a

View File

@ -13,7 +13,8 @@ class TestQuantized(mlx_tests.MLXTestCase):
w_q, scales, biases = mx.quantize(w, 64, b)
w_hat = mx.dequantize(w_q, scales, biases, 64, b)
errors = (w - w_hat).abs().reshape(*scales.shape, -1)
self.assertTrue((errors <= scales[..., None] / 2).all())
eps = 1e-6
self.assertTrue((errors <= (scales[..., None] / 2 + eps)).all())
def test_qmm(self):
key = mx.random.key(0)