(self)
| 1158 | ds = mx.grad(gmm)(s, x, wq) |
| 1159 | |
| 1160 | def test_quantize_strided(self): |
| 1161 | N = 64 |
| 1162 | mode = "nvfp4" |
| 1163 | w = mx.random.normal(shape=(N, N)) |
| 1164 | w_q, scales = mx.quantize(w, mode="nvfp4") |
| 1165 | |
| 1166 | scales = mx.broadcast_to(mx.array(56, mx.uint8), scales.shape) |
| 1167 | w_hat = mx.dequantize(w_q, scales, mode=mode) |
| 1168 | expected = mx.dequantize(w_q, mx.contiguous(scales), mode=mode) |
| 1169 | self.assertTrue(mx.allclose(w_hat, expected)) |
| 1170 | |
| 1171 | |
| 1172 | if __name__ == "__main__": |
nothing calls this directly
no test coverage detected