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 57a4f2f..cefc26b 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; @@ -275,17 +274,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 27fdcb9..d7a61d2 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", "