Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion python/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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:

Expand Down
4 changes: 4 additions & 0 deletions python/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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),
Expand Down
5 changes: 3 additions & 2 deletions python/src/scale.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
Loading