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
62 changes: 54 additions & 8 deletions python/src/input.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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";

Expand Down Expand Up @@ -39,9 +41,9 @@ pub fn as_f32_values<'py>(obj: &Bound<'py, PyAny>) -> PyResult<PyReadonlyArrayDy
read_f32(obj, "values", &[1], "a 1-D array").map(|(values, _)| values)
}

/// Anything that `numpy.asarray` turns into an array works, like a list or a
/// PyTorch tensor, as long as it holds real numbers and has one of
/// `dimensions`, which `wanted` describes. Errors call it `argument`. An
/// Anything that `numpy.asarray` turns into an array works, like a list, and
/// so does any PyTorch tensor, as long as it holds real numbers and has one
/// of `dimensions`, which `wanted` describes. Errors call it `argument`. An
/// array that already holds C-contiguous float32 values is read where it is,
/// without a copy.
fn read_f32<'py>(
Expand All @@ -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::<PyUntypedArray>()?;
if !matches!(dtype_kind(array)?.as_str(), "b" | "i" | "u" | "f") {
return Err(PyTypeError::new_err(format!(
Expand All @@ -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<Bound<'py, PyAny>> {
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<bool> {
if obj.is_instance_of::<PyUntypedArray>() {
return Ok(false);
}
let modules = obj.py().import("sys")?.getattr("modules")?;
match modules.cast::<PyDict>()?.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<String> {
Expand Down Expand Up @@ -153,10 +183,26 @@ pub fn as_i32_codes(obj: &Bound<'_, PyAny>) -> PyResult<Vec<i32>> {

/// Read packed codes: a 1-D uint8 array, like `Quantized.codes` returns.
pub fn as_packed_codes(obj: &Bound<'_, PyAny>) -> PyResult<Vec<u8>> {
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<Vec<u8>> {
match PyBuffer::<u8>::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<Vec<u8>> {
let array = obj.py().import("numpy")?.call_method1("asarray", (obj,))?;
let bytes = array
.cast::<PyArray1<u8>>()
.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`.
Expand Down
18 changes: 9 additions & 9 deletions python/src/quantized/methods.rs
Original file line number Diff line number Diff line change
@@ -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};
Expand All @@ -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;

Expand Down Expand Up @@ -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<u8>) -> PyResult<Self> {
let inner = QuantizedInner::from_bytes(&data.to_vec(py)?).map_err(from_quantize)?;
fn from_bytes(data: Bound<'_, PyAny>) -> PyResult<Self> {
let inner = QuantizedInner::from_bytes(&as_bytes(&data)?).map_err(from_quantize)?;
Ok(Self { inner })
}

#[new]
fn new(py: Python<'_>, data: PyBuffer<u8>) -> PyResult<Self> {
Self::from_bytes(py, data)
fn new(data: Bound<'_, PyAny>) -> PyResult<Self> {
Self::from_bytes(data)
}

fn __repr__(&self, py: Python<'_>) -> PyResult<String> {
Expand Down
59 changes: 59 additions & 0 deletions python/tests/test_quantize.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,8 @@
import io
import pickle
import re
import sys
import types
from collections.abc import Hashable

import numpy as np
Expand Down Expand Up @@ -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", "<U3")]:
with pytest.raises(TypeError, match=f"real numbers, got .* dtype {re.escape(dtype)}"):
Expand Down
Loading