From 09bd72d6f1dd730ab0e5345a3aa8698b8bfd9f51 Mon Sep 17 00:00:00 2001 From: Akshey D <131929364+aksheyd@users.noreply.github.com> Date: Mon, 28 Sep 2026 04:25:13 +0000 Subject: [PATCH 1/3] fix: pickle Quantized through its class, so torch.load can allow it torch.load defaults to weights_only=True, which refuses any global it doesn't trust. Pickles called the static method from_bytes, which pickles as builtins.getattr, so add_safe_globals([Quantized]) didn't help, and allowing getattr would defeat weights_only. Quantized(data) now loads the bytes that to_bytes saved, like from_bytes, and pickles call it, so allowing the class is enough. Pickles made through from_bytes still load, since from_bytes stays. --- python/README.md | 2 +- python/src/quantized/inner.rs | 4 +++ python/src/quantized/methods.rs | 17 ++++++++----- python/tests/test_quantize.py | 43 +++++++++++++++++++++++++++++++-- 4 files changed, 57 insertions(+), 9 deletions(-) diff --git a/python/README.md b/python/README.md index d4f1326..b443ae7 100644 --- a/python/README.md +++ b/python/README.md @@ -30,7 +30,7 @@ the other schemes return the same `Quantized` type: - `learned.alternate(q, weights)` refits too, then rounds each value to the nearest code on its block's new line, and repeats until no code moves. it also changes `q` in place - `Scheme.Q4_32.quantize(weights)` picks a scheme at run time -quantized values can be pickled, and compared with `==`. `q.to_bytes()` saves one as bytes, in the same format as the rust crate, and `Quantized.from_bytes(data)` loads it back. to keep it in an `np.savez` or safetensors file, store `np.frombuffer(q.to_bytes(), np.uint8)`. +quantized values can be pickled, and compared with `==`. to load them from a `torch.save` checkpoint, call `torch.serialization.add_safe_globals([Quantized])` before `torch.load`. `q.to_bytes()` saves one as bytes, in the same format as the rust crate, and `Quantized.from_bytes(data)` loads it back. to keep it in an `np.savez` or safetensors file, store `np.frombuffer(q.to_bytes(), np.uint8)`. to save its parts as plain arrays instead, like with `np.savez`, pass them back by name to `Quantized.from_parts`. an adaptive tensor keeps `block_bits` instead of `bits`: diff --git a/python/src/quantized/inner.rs b/python/src/quantized/inner.rs index 6a6ddff..a6fb859 100644 --- a/python/src/quantized/inner.rs +++ b/python/src/quantized/inner.rs @@ -138,6 +138,10 @@ impl QuantizedInner { /// Scales can be negative: a symmetric block puts its value farthest from /// zero on the most negative code, even when that value is positive. /// Zero-points are rarely whole numbers. +/// +/// `Quantized(data)` loads a tensor that `to_bytes` saved, like `from_bytes`. +/// Pickles load through it, so `torch.load` accepts them once +/// `torch.serialization.add_safe_globals([Quantized])` allows the class. #[pyclass(name = "Quantized", module = "quantize", eq)] #[derive(PartialEq)] pub struct PyQuantized { diff --git a/python/src/quantized/methods.rs b/python/src/quantized/methods.rs index 0fb693c..70c0bac 100644 --- a/python/src/quantized/methods.rs +++ b/python/src/quantized/methods.rs @@ -2,7 +2,7 @@ use numpy::{IntoPyArray, PyArray1, PyArrayMethods}; use pyo3::buffer::PyBuffer; use pyo3::exceptions::PyValueError; use pyo3::prelude::*; -use pyo3::types::{PyBytes, PyTuple}; +use pyo3::types::{PyBytes, PyTuple, PyType}; use quantize::{Error, Quantized, Scale}; @@ -280,6 +280,11 @@ impl PyQuantized { Ok(Self { inner }) } + #[new] + fn new(py: Python<'_>, data: PyBuffer) -> PyResult { + Self::from_bytes(py, data) + } + fn __repr__(&self, py: Python<'_>) -> PyResult { let shape = self.shape(py)?.repr()?; Ok(match self.bits() { @@ -300,11 +305,11 @@ impl PyQuantized { #[classattr] const __hash__: Option> = None; - fn __reduce__<'py>( - slf: &Bound<'py, Self>, - ) -> PyResult<(Bound<'py, PyAny>, (Bound<'py, PyBytes>,))> { - let from_bytes = slf.getattr("from_bytes")?; - Ok((from_bytes, (slf.borrow().to_bytes(slf.py()),))) + // Pickles rebuild the tensor by calling the class, so that `torch.load` + // loads them once `add_safe_globals([Quantized])` allows it. A static + // method would pickle as a call to `getattr`, which `torch.load` refuses. + fn __reduce__<'py>(slf: &Bound<'py, Self>) -> (Bound<'py, PyType>, (Bound<'py, PyBytes>,)) { + (slf.get_type(), (slf.borrow().to_bytes(slf.py()),)) } /// Pickles saved by quantize-py 0.2 call this with one tuple, which starts diff --git a/python/tests/test_quantize.py b/python/tests/test_quantize.py index 95c1acc..04ec94a 100644 --- a/python/tests/test_quantize.py +++ b/python/tests/test_quantize.py @@ -1,4 +1,5 @@ import array +import codecs import io import pickle import re @@ -453,11 +454,33 @@ def test_bytes_round_trip_every_kind_and_scale_type_through_numpy(): assert Quantized.from_bytes(saved["layer"]) == quantized -def test_pickles_hold_the_bytes_that_from_bytes_loads(): +def test_pickles_call_the_class_with_the_bytes_that_to_bytes_saves(): quantized = quantize(weight_matrix(8, 32), bits=4, scale=Scale.F16) rebuild, (data,) = quantized.__reduce__() - assert rebuild == Quantized.from_bytes + assert rebuild is Quantized assert data == quantized.to_bytes() + assert Quantized(data) == Quantized.from_bytes(data) == quantized + + +class OnlyQuantizedUnpickler(pickle.Unpickler): + """Stands in for `torch.load`, which by default refuses any global it + doesn't trust, after `torch.serialization.add_safe_globals([Quantized])`. + It also trusts `_codecs.encode`, as `torch.load` does, since pickles below + protocol 3 store bytes with it.""" + + def find_class(self, module, name): + if (module, name) == ("quantize", "Quantized"): + return Quantized + if (module, name) == ("_codecs", "encode"): + return codecs.encode + raise pickle.UnpicklingError(f"{module}.{name} isn't allowed") + + +def test_pickles_load_when_only_the_class_is_allowed_as_in_torch_load(): + checkpoint = {"layer": quantize(weight_matrix(8, 32), bits=4, scale=Scale.F16)} + for protocol in range(pickle.HIGHEST_PROTOCOL + 1): + data = pickle.dumps(checkpoint, protocol) + assert OnlyQuantizedUnpickler(io.BytesIO(data)).load() == checkpoint def test_from_bytes_rejects_bytes_that_do_not_hold_a_tensor(): @@ -494,6 +517,22 @@ def test_a_pickle_from_0_2_says_how_to_move_the_tensor_over(): Quantized._from_pickle((2, "symmetric")) +# quantize([0.42, -0.10, 0.70, -0.50], bits=8, block=4), pickled through +# Quantized.from_bytes by a development build of quantize-py 0.3.0. +PICKLED_THROUGH_FROM_BYTES = ( + b"\x80\x04\x95t\x00\x00\x00\x00\x00\x00\x00\x8c\x08builtins\x94\x8c\x07getattr" + b"\x94\x93\x94\x8c\x08quantize\x94\x8c\tQuantized\x94\x93\x94\x8c\nfrom_bytes\x94" + b"\x86\x94R\x94C+QNTZ\x01\x00\x03f32\x08\x04\x00\x00\x00\x00\x00\x00\x00\x04\x00" + b"\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x0033\xb3\xbb\xb3\x12\x80[" + b"\x94\x85\x94R\x94." +) + + +def test_a_pickle_through_from_bytes_still_loads(): + weights = [0.42, -0.10, 0.70, -0.50] + assert pickle.loads(PICKLED_THROUGH_FROM_BYTES) == quantize(weights, bits=8, block=4) + + def test_quantized_compares_by_value(): weights = weight_matrix(4, 32) quantized = quantize(weights, bits=4) From a940acb30f6e50500da77aa732df62cb8e051117 Mon Sep 17 00:00:00 2001 From: Akshey D <131929364+aksheyd@users.noreply.github.com> Date: Mon, 28 Sep 2026 04:28:46 +0000 Subject: [PATCH 2/3] docs: give the torch.load line its own paragraph after the np.savez recipe --- python/README.md | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/python/README.md b/python/README.md index b443ae7..abcf1c1 100644 --- a/python/README.md +++ b/python/README.md @@ -30,7 +30,7 @@ the other schemes return the same `Quantized` type: - `learned.alternate(q, weights)` refits too, then rounds each value to the nearest code on its block's new line, and repeats until no code moves. it also changes `q` in place - `Scheme.Q4_32.quantize(weights)` picks a scheme at run time -quantized values can be pickled, and compared with `==`. to load them from a `torch.save` checkpoint, call `torch.serialization.add_safe_globals([Quantized])` before `torch.load`. `q.to_bytes()` saves one as bytes, in the same format as the rust crate, and `Quantized.from_bytes(data)` loads it back. to keep it in an `np.savez` or safetensors file, store `np.frombuffer(q.to_bytes(), np.uint8)`. +quantized values can be pickled, and compared with `==`. `q.to_bytes()` saves one as bytes, in the same format as the rust crate, and `Quantized.from_bytes(data)` loads it back. to keep it in an `np.savez` or safetensors file, store `np.frombuffer(q.to_bytes(), np.uint8)`. to save its parts as plain arrays instead, like with `np.savez`, pass them back by name to `Quantized.from_parts`. an adaptive tensor keeps `block_bits` instead of `bits`: @@ -40,6 +40,8 @@ np.savez("layer.npz", kind=q.kind, shape=q.shape, block=q.block, bits=q.bits, q = Quantized.from_parts(**np.load("layer.npz")) ``` +to load quantized values from a `torch.save` checkpoint, call `torch.serialization.add_safe_globals([Quantized])` before `torch.load`. + each value decodes as `code * scale`, or `(code - zero_point) * scale` with zero-points, using the scale and zero-point of its block. codes are signed and `bits` wide, and `q.codes` packs them low bits first. scales can be negative, since a symmetric block puts its value farthest from zero on the most negative code. `help(Quantized)` has the details. to build and test from a clone of the repo, with rust 1.88 or newer and [just](https://github.com/casey/just): From 726255abc23f20f1e7d5b8e08016a811a0f66a79 Mon Sep 17 00:00:00 2001 From: Akshey D <131929364+aksheyd@users.noreply.github.com> Date: Mon, 28 Sep 2026 04:32:47 +0000 Subject: [PATCH 3/3] feat: read PyTorch parameters, bf16 tensors, and uint8 tensors for bytes quantize(linear.weight) raised PyTorch's RuntimeError, since numpy.asarray refuses a tensor that requires grad, and a bf16 weight raised "Got unsupported ScalarType BFloat16", since NumPy has no bfloat16. A PyTorch tensor is now detached, and a floating-point one read as float32, which is exact for bf16, before numpy.asarray reads it. torch is looked up in sys.modules, never imported, and NumPy arrays skip the lookup, so they cost what they did. from_bytes, Quantized(data), and from_parts' codes also take any 1-D uint8 array that numpy.asarray reads, like the tensor that safetensors.torch loads. --- python/src/input.rs | 62 ++++++++++++++++++++++++++++----- python/src/quantized/methods.rs | 18 +++++----- python/tests/test_quantize.py | 59 +++++++++++++++++++++++++++++++ 3 files changed, 122 insertions(+), 17 deletions(-) diff --git a/python/src/input.rs b/python/src/input.rs index 398e16b..8525fd1 100644 --- a/python/src/input.rs +++ b/python/src/input.rs @@ -3,14 +3,16 @@ use numpy::{ PyArray1, PyArrayDyn, PyArrayMethods, PyReadonlyArrayDyn, PyUntypedArray, PyUntypedArrayMethods, }; +use pyo3::buffer::PyBuffer; use pyo3::exceptions::{PyOverflowError, PyTypeError, PyValueError}; use pyo3::prelude::*; -use pyo3::types::{PyBool, PyTuple}; +use pyo3::types::{PyBool, PyDict, PyTuple}; use crate::error::{InvalidBitsError, InvalidBlockError, length_mismatch}; const CODES_TYPE: &str = "codes must be a 1-D signed integer array or a sequence of int; packed Quantized.codes is uint8 and must not be passed here — use unpacked_codes"; const PACKED_CODES_TYPE: &str = "codes must be a 1-D uint8 array, like Quantized.codes"; +const BYTES_TYPE: &str = "data must be bytes, like to_bytes returns, or a 1-D uint8 array"; const OUT_TYPE: &str = "out must be a writable C-contiguous native-endian float32 array"; const OUT_CONTIG: &str = "out must be writable and C-contiguous"; @@ -39,9 +41,9 @@ pub fn as_f32_values<'py>(obj: &Bound<'py, PyAny>) -> PyResult( @@ -56,7 +58,7 @@ fn read_f32<'py>( "{argument} can't be a masked array, since its mask would be ignored; fill in the masked values first, like {argument}.filled(0)" ))); } - let converted = numpy.call_method1("asarray", (obj,))?; + let converted = numpy.call_method1("asarray", (readable_by_numpy(obj)?,))?; let array = converted.cast::()?; if !matches!(dtype_kind(array)?.as_str(), "b" | "i" | "u" | "f") { return Err(PyTypeError::new_err(format!( @@ -76,6 +78,34 @@ fn read_f32<'py>( Ok((typed.try_readonly()?, array.shape().to_vec())) } +/// `obj`, or if it's a PyTorch tensor, a tensor that `numpy.asarray` reads: +/// detached, since NumPy refuses one that requires grad, and in float32 if +/// it holds floating-point numbers, since NumPy has no bfloat16. +fn readable_by_numpy<'py>(obj: &Bound<'py, PyAny>) -> PyResult> { + if !is_torch_tensor(obj)? { + return Ok(obj.clone()); + } + let tensor = obj.call_method0("detach")?; + if tensor.call_method0("is_floating_point")?.is_truthy()? { + return tensor.call_method0("float"); + } + Ok(tensor) +} + +/// Whether `obj` is a PyTorch tensor. A program that passes one has imported +/// torch, so torch is looked up in `sys.modules` instead of imported, and +/// NumPy arrays, the usual input, skip even that. +fn is_torch_tensor(obj: &Bound<'_, PyAny>) -> PyResult { + if obj.is_instance_of::() { + return Ok(false); + } + let modules = obj.py().import("sys")?.getattr("modules")?; + match modules.cast::()?.get_item("torch")? { + Some(torch) if !torch.is_none() => obj.is_instance(&torch.getattr("Tensor")?), + _ => Ok(false), + } +} + /// `obj`'s type, and the shape and dtype that `numpy.asarray` gave it, such /// as `torch.Tensor with shape (2, 3, 4) and dtype float32`. fn describe(obj: &Bound<'_, PyAny>, array: &Bound<'_, PyUntypedArray>) -> PyResult { @@ -153,10 +183,26 @@ pub fn as_i32_codes(obj: &Bound<'_, PyAny>) -> PyResult> { /// Read packed codes: a 1-D uint8 array, like `Quantized.codes` returns. pub fn as_packed_codes(obj: &Bound<'_, PyAny>) -> PyResult> { - let codes = obj + read_uint8(obj, PACKED_CODES_TYPE) +} + +/// Read the bytes that `to_bytes` saved: `bytes` or another bytes-like +/// object, or a 1-D uint8 array. +pub fn as_bytes(obj: &Bound<'_, PyAny>) -> PyResult> { + match PyBuffer::::get(obj) { + Ok(buffer) => buffer.to_vec(obj.py()), + Err(_) => read_uint8(obj, BYTES_TYPE), + } +} + +/// Read a 1-D uint8 array, or anything that `numpy.asarray` turns into one, +/// like a PyTorch tensor. Anything else raises `TypeError(message)`. +fn read_uint8(obj: &Bound<'_, PyAny>, message: &'static str) -> PyResult> { + let array = obj.py().import("numpy")?.call_method1("asarray", (obj,))?; + let bytes = array .cast::>() - .map_err(|_| PyTypeError::new_err(PACKED_CODES_TYPE))?; - Ok(codes.try_readonly()?.as_array().to_vec()) + .map_err(|_| PyTypeError::new_err(message))?; + Ok(bytes.try_readonly()?.as_array().to_vec()) } /// Borrow `out` for writing, after checking it has exactly `shape`. diff --git a/python/src/quantized/methods.rs b/python/src/quantized/methods.rs index 70c0bac..bee0d2e 100644 --- a/python/src/quantized/methods.rs +++ b/python/src/quantized/methods.rs @@ -1,5 +1,4 @@ use numpy::{IntoPyArray, PyArray1, PyArrayMethods}; -use pyo3::buffer::PyBuffer; use pyo3::exceptions::PyValueError; use pyo3::prelude::*; use pyo3::types::{PyBytes, PyTuple, PyType}; @@ -10,8 +9,8 @@ use super::inner::{PyQuantized, QuantizedInner, with_inner}; use super::parts::Parts; use crate::error::from_quantize; use crate::input::{ - as_f32_array, as_f32_matmul_values, as_f32_values, as_packed_codes, as_writable_f32_out, - check_shape, + as_bytes, as_f32_array, as_f32_matmul_values, as_f32_values, as_packed_codes, + as_writable_f32_out, check_shape, }; use crate::scale::PyScale; @@ -272,17 +271,18 @@ impl PyQuantized { } /// Load a tensor that `to_bytes` saved, in Python or in Rust. `data` is - /// `bytes` or another bytes-like object, such as a uint8 NumPy array. - /// Bytes that don't hold a valid tensor raise `QuantizeError`. + /// `bytes` or another bytes-like object, or a 1-D uint8 array, such as a + /// NumPy array or the PyTorch tensor that safetensors loads. Bytes that + /// don't hold a valid tensor raise `QuantizeError`. #[staticmethod] - fn from_bytes(py: Python<'_>, data: PyBuffer) -> PyResult { - let inner = QuantizedInner::from_bytes(&data.to_vec(py)?).map_err(from_quantize)?; + fn from_bytes(data: Bound<'_, PyAny>) -> PyResult { + let inner = QuantizedInner::from_bytes(&as_bytes(&data)?).map_err(from_quantize)?; Ok(Self { inner }) } #[new] - fn new(py: Python<'_>, data: PyBuffer) -> PyResult { - Self::from_bytes(py, data) + fn new(data: Bound<'_, PyAny>) -> PyResult { + Self::from_bytes(data) } fn __repr__(&self, py: Python<'_>) -> PyResult { diff --git a/python/tests/test_quantize.py b/python/tests/test_quantize.py index 04ec94a..4df3b73 100644 --- a/python/tests/test_quantize.py +++ b/python/tests/test_quantize.py @@ -3,6 +3,8 @@ import io import pickle import re +import sys +import types from collections.abc import Hashable import numpy as np @@ -179,6 +181,63 @@ def test_quantize_reads_anything_numpy_asarray_reads(): assert quantize(Tensor(weights)).matmul(Tensor(weights[:2])).shape == (2, 4) +class TorchTensor(Tensor): + """Stands in for a PyTorch tensor, which NumPy can't read while it + requires grad or holds bfloat16.""" + + def __init__(self, values, requires_grad=False, dtype="float32"): + super().__init__(values) + self.requires_grad = requires_grad + self.dtype = dtype + + def __array__(self, dtype=None, copy=None): + if self.requires_grad: + raise RuntimeError("Can't call numpy() on Tensor that requires grad.") + if self.dtype == "bfloat16": + raise TypeError("Got unsupported ScalarType BFloat16") + return self.values + + def detach(self): + return TorchTensor(self.values, dtype=self.dtype) + + def is_floating_point(self): + return self.values.dtype.kind == "f" + + def float(self): + return TorchTensor(self.values, self.requires_grad) + + +def test_torch_tensors_are_read_even_if_they_require_grad_or_hold_bfloat16(monkeypatch): + monkeypatch.setitem(sys.modules, "torch", types.SimpleNamespace(Tensor=TorchTensor)) + weights = weight_matrix(4, 32) + expected = quantize(weights) + for tensor in [TorchTensor(weights, requires_grad=True), TorchTensor(weights, dtype="bfloat16")]: + assert quantize(tensor) == expected + assert expected.dot(tensor) == expected.dot(weights) + # Some programs keep torch from being imported this way. + monkeypatch.setitem(sys.modules, "torch", None) + assert quantize(weights.tolist()) == expected + + +def test_bytes_and_codes_can_be_any_uint8_array_that_numpy_asarray_reads(): + quantized = quantize(weight_matrix(4, 32), bits=4) + data = Tensor(np.frombuffer(quantized.to_bytes(), np.uint8)) + assert Quantized.from_bytes(data) == Quantized(data) == quantized + rebuilt = Quantized.from_parts( + kind="symmetric", + shape=(4, 32), + block=32, + bits=4, + codes=Tensor(quantized.codes), + scales=quantized.scales, + scale="f32", + ) + assert rebuilt == quantized + for wrong in [Tensor(quantized.unpacked_codes), "QNTZ", [81, 78, 84, 90]]: + with pytest.raises(TypeError, match="data must be bytes, like to_bytes returns, or a 1-D uint8"): + Quantized.from_bytes(wrong) + + def test_values_that_are_not_real_numbers_are_rejected_with_what_arrived(): for values, dtype in [(np.array([1 + 2j]), "complex128"), (b"ab", "|S2"), ("0.5", "