From 9d5e4b27eeb4cc575d04e5aaf5903efba8287046 Mon Sep 17 00:00:00 2001 From: Akshey D <131929364+aksheyd@users.noreply.github.com> Date: Mon, 28 Sep 2026 04:34:56 +0000 Subject: [PATCH] docs: recommend f32 scales for asymmetric codes above about 10 bits f16 and bf16 store each block's zero-point too, and their rounding grows with the bit width, while the README and docstrings offered 2 to 16 bits with any scale type. On normal weights, asymmetric bf16's worst error is 2.3 times f32's at 11 bits and 71 times at 16, and f16's is 7 times at 16. Symmetric codes are unaffected. The README, asymmetric.quantize, and Scale now say to use f32 scales there. --- python/README.md | 2 +- python/src/lib.rs | 4 ++++ python/src/scale.rs | 5 +++-- 3 files changed, 8 insertions(+), 3 deletions(-) diff --git a/python/README.md b/python/README.md index d4f1326..edc3523 100644 --- a/python/README.md +++ b/python/README.md @@ -20,7 +20,7 @@ dot = q.dot(weights) # 0.926 `bits` is the width of each code, from 2 to 16. `block` is how many values share one scale, and `scale` is how that scale is stored: `Scale.F32` (the default), `Scale.F16`, or `Scale.BF16`. values can be a list, a numpy array, or anything else `np.asarray` reads, like a pytorch tensor. a 2-d array keeps its shape, so `q.dequantize()` gives back a matrix and `q.matmul(x)` computes `x @ W.T`, like a linear layer. -the scales count toward the size: 4-bit codes with one f16 scale per 32 values cost 4.5 bits per value, or 5 with the default f32 scale. `q.bits_per_element` reports it. +the scales count toward the size: 4-bit codes with one f16 scale per 32 values cost 4.5 bits per value, or 5 with the default f32 scale. `q.bits_per_element` reports it. above about 10 bits, use f32 scales with `asymmetric.quantize`, since f16 and bf16 zero-points cap its accuracy. the other schemes return the same `Quantized` type: diff --git a/python/src/lib.rs b/python/src/lib.rs index e0c282a..9898c13 100644 --- a/python/src/lib.rs +++ b/python/src/lib.rs @@ -80,6 +80,10 @@ fn quantize_tensor( /// scale and a zero-point, so values that aren't centered on zero can use /// every code. Each value decodes as `(code - zero_point) * scale`. The other /// arguments work as in `quantize`. +/// +/// `Scale.F16` and `Scale.BF16` round each zero-point, which caps the +/// accuracy above about 10 bits, or sooner on blocks far from zero, so use +/// `Scale.F32` there. #[pyfunction] #[pyo3( signature = (values, bits = 8, block = 32, *, scale = PyScale::F32), diff --git a/python/src/scale.rs b/python/src/scale.rs index f162491..940b099 100644 --- a/python/src/scale.rs +++ b/python/src/scale.rs @@ -7,8 +7,9 @@ use crate::error::QuantizeError; /// How each block's scale and zero-point are stored. `Scale.F32` keeps them /// exactly, in 4 bytes each. `Scale.F16` and `Scale.BF16` round them to 2 -/// bytes: f16 keeps more digits, and bf16 more range. `scale=` also takes -/// the name that `name` returns. +/// bytes: f16 keeps more digits, and bf16 more range. Rounded zero-points cap +/// the accuracy of asymmetric codes above about 10 bits, so use `Scale.F32` +/// there. `scale=` also takes the name that `name` returns. #[pyclass( eq, frozen,