quantize-py 0.3.0__tar.gz → 0.3.1__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (46) hide show
  1. {quantize_py-0.3.0 → quantize_py-0.3.1}/Cargo.lock +2 -2
  2. {quantize_py-0.3.0 → quantize_py-0.3.1}/Cargo.toml +3 -6
  3. {quantize_py-0.3.0 → quantize_py-0.3.1}/PKG-INFO +8 -4
  4. {quantize_py-0.3.0 → quantize_py-0.3.1}/README.md +10 -4
  5. {quantize_py-0.3.0 → quantize_py-0.3.1}/python/README.md +7 -3
  6. {quantize_py-0.3.0 → quantize_py-0.3.1}/python/src/lib.rs +5 -1
  7. {quantize_py-0.3.0 → quantize_py-0.3.1}/python/src/quantized/inner.rs +31 -14
  8. {quantize_py-0.3.0 → quantize_py-0.3.1}/python/src/quantized/methods.rs +10 -6
  9. {quantize_py-0.3.0 → quantize_py-0.3.1}/python/src/scale.rs +15 -4
  10. {quantize_py-0.3.0 → quantize_py-0.3.1}/python/src/scheme.rs +25 -3
  11. {quantize_py-0.3.0 → quantize_py-0.3.1}/python/tests/test_learned.py +10 -0
  12. {quantize_py-0.3.0 → quantize_py-0.3.1}/python/tests/test_quantize.py +56 -7
  13. {quantize_py-0.3.0 → quantize_py-0.3.1}/src/kernels/block.rs +38 -37
  14. {quantize_py-0.3.0 → quantize_py-0.3.1}/src/kernels/i4.rs +14 -4
  15. {quantize_py-0.3.0 → quantize_py-0.3.1}/src/kernels/i8.rs +14 -4
  16. {quantize_py-0.3.0 → quantize_py-0.3.1}/src/lib.rs +7 -6
  17. {quantize_py-0.3.0 → quantize_py-0.3.1}/src/methods/adaptive/mod.rs +2 -4
  18. {quantize_py-0.3.0 → quantize_py-0.3.1}/src/methods/learned/mod.rs +29 -1
  19. {quantize_py-0.3.0 → quantize_py-0.3.1}/src/methods/symmetric/mod.rs +50 -0
  20. {quantize_py-0.3.0 → quantize_py-0.3.1}/src/shared/decode.rs +187 -150
  21. {quantize_py-0.3.0 → quantize_py-0.3.1}/src/shared/error.rs +67 -5
  22. {quantize_py-0.3.0 → quantize_py-0.3.1}/src/shared/packed.rs +11 -13
  23. quantize_py-0.3.1/src/shared/scheme.rs +252 -0
  24. {quantize_py-0.3.0 → quantize_py-0.3.1}/src/shared/tensor.rs +83 -27
  25. quantize_py-0.3.0/src/shared/scheme.rs +0 -69
  26. {quantize_py-0.3.0 → quantize_py-0.3.1}/LICENSE +0 -0
  27. {quantize_py-0.3.0 → quantize_py-0.3.1}/pyproject.toml +0 -0
  28. {quantize_py-0.3.0 → quantize_py-0.3.1}/python/Cargo.toml +0 -0
  29. {quantize_py-0.3.0 → quantize_py-0.3.1}/python/quantize/__init__.py +0 -0
  30. {quantize_py-0.3.0 → quantize_py-0.3.1}/python/quantize/adaptive.py +0 -0
  31. {quantize_py-0.3.0 → quantize_py-0.3.1}/python/quantize/asymmetric.py +0 -0
  32. {quantize_py-0.3.0 → quantize_py-0.3.1}/python/quantize/learned.py +0 -0
  33. {quantize_py-0.3.0 → quantize_py-0.3.1}/python/quantize/symmetric.py +0 -0
  34. {quantize_py-0.3.0 → quantize_py-0.3.1}/python/src/error.rs +0 -0
  35. {quantize_py-0.3.0 → quantize_py-0.3.1}/python/src/input.rs +0 -0
  36. {quantize_py-0.3.0 → quantize_py-0.3.1}/python/src/learned.rs +0 -0
  37. {quantize_py-0.3.0 → quantize_py-0.3.1}/python/src/quantized/parts.rs +0 -0
  38. {quantize_py-0.3.0 → quantize_py-0.3.1}/python/src/quantized.rs +0 -0
  39. {quantize_py-0.3.0 → quantize_py-0.3.1}/src/kernels/mod.rs +0 -0
  40. {quantize_py-0.3.0 → quantize_py-0.3.1}/src/kernels/reduce.rs +0 -0
  41. {quantize_py-0.3.0 → quantize_py-0.3.1}/src/methods/asymmetric/mod.rs +0 -0
  42. {quantize_py-0.3.0 → quantize_py-0.3.1}/src/methods/mod.rs +0 -0
  43. {quantize_py-0.3.0 → quantize_py-0.3.1}/src/shared/bytes.rs +0 -0
  44. {quantize_py-0.3.0 → quantize_py-0.3.1}/src/shared/mod.rs +0 -0
  45. {quantize_py-0.3.0 → quantize_py-0.3.1}/src/shared/params.rs +0 -0
  46. {quantize_py-0.3.0 → quantize_py-0.3.1}/src/shared/scale.rs +0 -0
@@ -1411,14 +1411,14 @@ dependencies = [
1411
1411
 
1412
1412
  [[package]]
1413
1413
  name = "quantize"
1414
- version = "0.3.0"
1414
+ version = "0.3.1"
1415
1415
  dependencies = [
1416
1416
  "half",
1417
1417
  ]
1418
1418
 
1419
1419
  [[package]]
1420
1420
  name = "quantize-py"
1421
- version = "0.3.0"
1421
+ version = "0.3.1"
1422
1422
  dependencies = [
1423
1423
  "half",
1424
1424
  "numpy",
@@ -17,16 +17,13 @@ include = [
17
17
  "/README.md",
18
18
  ]
19
19
 
20
- [package.metadata.docs.rs]
21
- features = ["std"]
22
-
23
20
  [workspace]
24
21
  members = ["python"]
25
22
  resolver = "3"
26
23
 
27
24
  [workspace.package]
28
25
  # Shared with the Python package. Between releases it ends in -dev.
29
- version = "0.3.0"
26
+ version = "0.3.1"
30
27
 
31
28
  [workspace.dependencies]
32
29
  half = "2"
@@ -38,8 +35,8 @@ all = { level = "deny", priority = -1 }
38
35
 
39
36
  [features]
40
37
  default = ["std"]
41
- # Reserved for future no_std support. For now it only implements
42
- # std::error::Error for Error.
38
+ # Does nothing, since the crate always needs the standard library. Kept so
39
+ # that a Cargo.toml that names it still builds.
43
40
  std = []
44
41
 
45
42
  [dependencies]
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: quantize-py
3
- Version: 0.3.0
3
+ Version: 0.3.1
4
4
  Classifier: Programming Language :: Python :: 3
5
5
  Classifier: Programming Language :: Python :: 3.12
6
6
  Classifier: Programming Language :: Python :: 3.13
@@ -27,7 +27,7 @@ python bindings for [quantize](https://github.com/aksheyd/quantize), a simple, f
27
27
  pip install quantize-py
28
28
  ```
29
29
 
30
- requires python 3.12 or newer. numpy is installed with it.
30
+ requires python 3.12 or newer. numpy is installed with it. where no wheel fits, like free-threaded python or alpine linux, pip builds it from source, which needs rust 1.88 or newer.
31
31
 
32
32
  ```python
33
33
  from quantize import Scale, quantize
@@ -48,12 +48,16 @@ the other schemes return the same `Quantized` type:
48
48
  - `asymmetric.quantize(weights, bits=8, block=32)` adds a zero-point per block, for values that aren't centered on zero
49
49
  - `adaptive.quantize(weights, tolerance=0.1 * weights.std())` gives each block the fewest bits, from 2 to 8, that round every weight within `tolerance`, in the weights' own units. a tenth of their standard deviation gives about 5 bits a block. for a list, use `np.std(weights)`
50
50
  - `learned.refine(q, weights)` refits each block's scale, and its zero-point if it has one, to lower the mean squared error. it changes `q` in place, so call `q.copy()` first to keep the original
51
- - `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. both can raise the worst error past an adaptive tensor's tolerance
51
+ - `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. both can raise the worst error past an adaptive tensor's tolerance, and lowering the error of the weights doesn't always lower the error of a model's outputs, so check those too
52
52
  - `Scheme.Q4_32.quantize(weights)` picks a scheme at run time
53
53
 
54
54
  a block with outliers can need more than 8 bits, which raises `ToleranceTooTightError`. retrying with its `smallest_tolerance` works, but loosens every block, not just that one:
55
55
 
56
56
  ```python
57
+ import numpy as np
58
+ from quantize import ToleranceTooTightError, adaptive
59
+
60
+ weights = np.array(weights) # a list has no .std(). skip this line for a pytorch tensor
57
61
  try:
58
62
  q = adaptive.quantize(weights, tolerance=0.1 * weights.std())
59
63
  except ToleranceTooTightError as error:
@@ -71,7 +75,7 @@ np.savez("layer.npz", **{name: part for name, part in parts.items() if part is n
71
75
  q = Quantized.from_parts(**np.load("layer.npz"))
72
76
  ```
73
77
 
74
- to load quantized values from a `torch.save` checkpoint, call `torch.serialization.add_safe_globals([Quantized])` before `torch.load`.
78
+ to keep quantized values in a `torch.save` checkpoint, store `torch.frombuffer(bytearray(q.to_bytes()), dtype=torch.uint8)`, and load each back with `Quantized(t)`. the checkpoint is then as small as the bytes, and `torch.load` reads it without `add_safe_globals`. pickling `q` itself makes the checkpoint about 1.5 times larger, since `torch.save` stores bytes as text, and needs `torch.serialization.add_safe_globals([Quantized])` before `torch.load`.
75
79
 
76
80
  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.
77
81
 
@@ -59,7 +59,7 @@ against [candle](https://github.com/huggingface/candle)'s `Q4_0`, `Q5_0`, and `Q
59
59
  ### quality
60
60
 
61
61
  ```
62
- cargo run --release --example compare
62
+ cargo run --release -p benchmarks --example compare
63
63
  ```
64
64
 
65
65
  quantize two random matrices, reconstruct them, then matmul. mse is the mean squared error against the exact f32 result: smaller is better, and it varies a little between runs. bits/value counts the scales too: 4-bit codes plus one f16 scale per 32 values is 4 + 16/32 = 4.5.
@@ -74,15 +74,21 @@ quantize two random matrices, reconstruct them, then matmul. mse is the mean squ
74
74
 
75
75
  <!-- comparison:end -->
76
76
 
77
- on WikiText-2 with SmolLM-135M, 4-bit perplexity is 26.44 against candle `Q4_0`'s 26.46, and fp32 is 18.93. lower is better, and `cargo run --release --example wikitext --features benchmarks/workload` reproduces them.
77
+ on WikiText-2 with SmolLM-135M, 4-bit perplexity is 26.44 against candle `Q4_0`'s 26.46, and fp32 is 18.93. perplexity is roughly how many tokens the model is choosing between when it guesses the next one, so lower is better. reproducing these numbers downloads the model and dataset from hugging face, then scores the whole test set, which takes hours:
78
+
79
+ ```
80
+ cargo run --release -p benchmarks --example wikitext --features workload
81
+ ```
82
+
83
+ add `-- --max-tokens 512` to score only the first 512 tokens. that takes about a minute, though its numbers won't match.
78
84
 
79
85
  ### speed
80
86
 
81
87
  ```
82
- cargo run --release --example throughput
88
+ cargo run --release -p benchmarks --example throughput
83
89
  ```
84
90
 
85
- quantize and dequantize with f16 scales. both libraries allocate their output on every call. aarch64 is an apple M5 Max and x86_64 an intel xeon. the hand-written simd only targets aarch64, so x86_64 runs plain loops and is slower.
91
+ quantize and dequantize with f16 scales. both libraries allocate their output on every call. aarch64 is an apple M5 Max and x86_64 an intel xeon. the hand-written simd only targets aarch64, so x86_64 runs plain loops and is slower. the run also times this crate's `dot`, `matmul` on 16 inputs, and row-by-row decoding of an adaptive matrix, which aren't compared with candle.
86
92
 
87
93
  <!-- speed:start -->
88
94
 
@@ -6,7 +6,7 @@ python bindings for [quantize](https://github.com/aksheyd/quantize), a simple, f
6
6
  pip install quantize-py
7
7
  ```
8
8
 
9
- requires python 3.12 or newer. numpy is installed with it.
9
+ requires python 3.12 or newer. numpy is installed with it. where no wheel fits, like free-threaded python or alpine linux, pip builds it from source, which needs rust 1.88 or newer.
10
10
 
11
11
  ```python
12
12
  from quantize import Scale, quantize
@@ -27,12 +27,16 @@ the other schemes return the same `Quantized` type:
27
27
  - `asymmetric.quantize(weights, bits=8, block=32)` adds a zero-point per block, for values that aren't centered on zero
28
28
  - `adaptive.quantize(weights, tolerance=0.1 * weights.std())` gives each block the fewest bits, from 2 to 8, that round every weight within `tolerance`, in the weights' own units. a tenth of their standard deviation gives about 5 bits a block. for a list, use `np.std(weights)`
29
29
  - `learned.refine(q, weights)` refits each block's scale, and its zero-point if it has one, to lower the mean squared error. it changes `q` in place, so call `q.copy()` first to keep the original
30
- - `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. both can raise the worst error past an adaptive tensor's tolerance
30
+ - `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. both can raise the worst error past an adaptive tensor's tolerance, and lowering the error of the weights doesn't always lower the error of a model's outputs, so check those too
31
31
  - `Scheme.Q4_32.quantize(weights)` picks a scheme at run time
32
32
 
33
33
  a block with outliers can need more than 8 bits, which raises `ToleranceTooTightError`. retrying with its `smallest_tolerance` works, but loosens every block, not just that one:
34
34
 
35
35
  ```python
36
+ import numpy as np
37
+ from quantize import ToleranceTooTightError, adaptive
38
+
39
+ weights = np.array(weights) # a list has no .std(). skip this line for a pytorch tensor
36
40
  try:
37
41
  q = adaptive.quantize(weights, tolerance=0.1 * weights.std())
38
42
  except ToleranceTooTightError as error:
@@ -50,7 +54,7 @@ np.savez("layer.npz", **{name: part for name, part in parts.items() if part is n
50
54
  q = Quantized.from_parts(**np.load("layer.npz"))
51
55
  ```
52
56
 
53
- to load quantized values from a `torch.save` checkpoint, call `torch.serialization.add_safe_globals([Quantized])` before `torch.load`.
57
+ to keep quantized values in a `torch.save` checkpoint, store `torch.frombuffer(bytearray(q.to_bytes()), dtype=torch.uint8)`, and load each back with `Quantized(t)`. the checkpoint is then as small as the bytes, and `torch.load` reads it without `add_safe_globals`. pickling `q` itself makes the checkpoint about 1.5 times larger, since `torch.save` stores bytes as text, and needs `torch.serialization.add_safe_globals([Quantized])` before `torch.load`.
54
58
 
55
59
  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.
56
60
 
@@ -147,7 +147,11 @@ fn adaptive_quantize(
147
147
  quantize_values(py, values, scale, |_| Scheme::Adaptive { block, tolerance })
148
148
  }
149
149
 
150
- #[pymodule]
150
+ // `refine` and `alternate` borrow a tensor mutably while they change it, and
151
+ // only the GIL keeps other threads from borrowing it at the same time, which
152
+ // would panic. So free-threaded Python turns the GIL back on when it imports
153
+ // this module.
154
+ #[pymodule(gil_used = true)]
151
155
  #[pyo3(name = "_native")]
152
156
  fn native(m: &Bound<'_, PyModule>) -> PyResult<()> {
153
157
  m.add("__version__", env!("CARGO_PKG_VERSION"))?;
@@ -1,3 +1,5 @@
1
+ use std::sync::Arc;
2
+
1
3
  use half::{bf16, f16};
2
4
  use pyo3::prelude::*;
3
5
  use quantize::{Quantized, Scale, Scheme, learned};
@@ -6,11 +8,16 @@ use super::parts::Parts;
6
8
  use crate::error::from_quantize;
7
9
  use crate::scale::PyScale;
8
10
 
11
+ /// The tensor, with one of three scale types, behind an `Arc` so that a
12
+ /// clone shares its codes and scales instead of copying them. `refine` and
13
+ /// `alternate` change it through `Arc::make_mut`, which copies a shared
14
+ /// tensor first, so a `copy()`, or a `dot` or `matmul` still running, keeps
15
+ /// the values it had.
9
16
  #[derive(Clone, PartialEq)]
10
17
  pub(crate) enum QuantizedInner {
11
- F32(Quantized<f32>),
12
- F16(Quantized<f16>),
13
- Bf16(Quantized<bf16>),
18
+ F32(Arc<Quantized<f32>>),
19
+ F16(Arc<Quantized<f16>>),
20
+ Bf16(Arc<Quantized<bf16>>),
14
21
  }
15
22
 
16
23
  macro_rules! with_inner {
@@ -64,9 +71,15 @@ impl QuantizedInner {
64
71
  scale: PyScale,
65
72
  ) -> PyResult<Self> {
66
73
  match scale {
67
- PyScale::F32 => quantize_shaped(scheme, values, shape).map(Self::F32),
68
- PyScale::F16 => quantize_shaped(scheme, values, shape).map(Self::F16),
69
- PyScale::Bf16 => quantize_shaped(scheme, values, shape).map(Self::Bf16),
74
+ PyScale::F32 => quantize_shaped(scheme, values, shape)
75
+ .map(Arc::new)
76
+ .map(Self::F32),
77
+ PyScale::F16 => quantize_shaped(scheme, values, shape)
78
+ .map(Arc::new)
79
+ .map(Self::F16),
80
+ PyScale::Bf16 => quantize_shaped(scheme, values, shape)
81
+ .map(Arc::new)
82
+ .map(Self::Bf16),
70
83
  }
71
84
  .map_err(from_quantize)
72
85
  }
@@ -92,11 +105,15 @@ impl QuantizedInner {
92
105
  }
93
106
 
94
107
  pub(crate) fn refine(&mut self, values: &[f32]) -> quantize::Result<()> {
95
- with_inner!(self, |quantized| learned::refine(quantized, values))
108
+ with_inner!(self, |quantized| {
109
+ learned::refine(Arc::make_mut(quantized), values)
110
+ })
96
111
  }
97
112
 
98
113
  pub(crate) fn alternate(&mut self, values: &[f32]) -> quantize::Result<bool> {
99
- with_inner!(self, |quantized| learned::alternate(quantized, values))
114
+ with_inner!(self, |quantized| {
115
+ learned::alternate(Arc::make_mut(quantized), values)
116
+ })
100
117
  }
101
118
 
102
119
  pub(crate) fn to_bytes(&self) -> Vec<u8> {
@@ -107,17 +124,17 @@ impl QuantizedInner {
107
124
  /// their header names.
108
125
  pub(crate) fn from_bytes(bytes: &[u8]) -> quantize::Result<Self> {
109
126
  match saved_scale(bytes) {
110
- PyScale::F32 => Quantized::from_bytes(bytes).map(Self::F32),
111
- PyScale::F16 => Quantized::from_bytes(bytes).map(Self::F16),
112
- PyScale::Bf16 => Quantized::from_bytes(bytes).map(Self::Bf16),
127
+ PyScale::F32 => Quantized::from_bytes(bytes).map(Arc::new).map(Self::F32),
128
+ PyScale::F16 => Quantized::from_bytes(bytes).map(Arc::new).map(Self::F16),
129
+ PyScale::Bf16 => Quantized::from_bytes(bytes).map(Arc::new).map(Self::Bf16),
113
130
  }
114
131
  }
115
132
 
116
133
  pub(crate) fn from_parts(parts: Parts, scale: PyScale) -> PyResult<Self> {
117
134
  match scale {
118
- PyScale::F32 => parts.into_quantized().map(Self::F32),
119
- PyScale::F16 => parts.into_quantized().map(Self::F16),
120
- PyScale::Bf16 => parts.into_quantized().map(Self::Bf16),
135
+ PyScale::F32 => parts.into_quantized().map(Arc::new).map(Self::F32),
136
+ PyScale::F16 => parts.into_quantized().map(Arc::new).map(Self::F16),
137
+ PyScale::Bf16 => parts.into_quantized().map(Arc::new).map(Self::Bf16),
121
138
  }
122
139
  }
123
140
  }
@@ -39,9 +39,10 @@ fn bits<S: Scale>(quantized: &Quantized<S>) -> Option<u32> {
39
39
  }
40
40
  }
41
41
 
42
- // Methods that release the GIL take `slf` instead of `&self`, and clone the
43
- // tensor out of a short borrow first. A borrow held while the GIL is released
44
- // would keep `learned.refine` on another thread from borrowing it mutably.
42
+ // Methods that release the GIL take `slf` instead of `&self`, and clone
43
+ // `inner` out of a short borrow first, which shares the tensor instead of
44
+ // copying it. A borrow held while the GIL is released would keep
45
+ // `learned.refine` on another thread from borrowing it mutably.
45
46
  #[pymethods]
46
47
  impl PyQuantized {
47
48
  /// Decode the values into an array of the tensor's `shape`. `out`, if
@@ -72,9 +73,8 @@ impl PyQuantized {
72
73
  /// The dot product of the decoded values with `values`, an array of the
73
74
  /// tensor's `shape`. For a matrix times a vector, use `matmul`.
74
75
  ///
75
- /// Symmetric tensors with 4-bit or 8-bit codes are decoded straight from
76
- /// the packed codes, without storing the decoded values. Other tensors are
77
- /// first unpacked into a buffer as large as the decoded values.
76
+ /// Each code is read straight from the packed bytes as it's multiplied,
77
+ /// so the decoded values are never stored, whatever the scheme.
78
78
  fn dot(slf: &Bound<'_, Self>, values: Bound<'_, PyAny>) -> PyResult<f32> {
79
79
  let (array, values_shape) = as_f32_array(&values)?;
80
80
  let values = array.as_slice()?;
@@ -93,6 +93,10 @@ impl PyQuantized {
93
93
  /// `inputs` is one vector of shape `(columns,)` or a batch of shape
94
94
  /// `(batch, columns)`. The result is `inputs @ W.T`, of shape `(rows,)`
95
95
  /// or `(batch, rows)`.
96
+ ///
97
+ /// Each call decodes the matrix one row at a time, straight from the
98
+ /// packed codes, and multiplies each row by every input before moving
99
+ /// on, so the whole matrix is never decoded at once.
96
100
  fn matmul<'py>(
97
101
  slf: &Bound<'py, Self>,
98
102
  inputs: Bound<'_, PyAny>,
@@ -2,6 +2,7 @@
2
2
 
3
3
  use pyo3::exceptions::PyValueError;
4
4
  use pyo3::prelude::*;
5
+ use pyo3::types::PyType;
5
6
 
6
7
  use crate::error::QuantizeError;
7
8
 
@@ -9,7 +10,10 @@ use crate::error::QuantizeError;
9
10
  /// exactly, in 4 bytes each. `Scale.F16` and `Scale.BF16` round them to 2
10
11
  /// bytes: f16 keeps more digits, and bf16 more range. Rounded zero-points cap
11
12
  /// the accuracy of asymmetric codes above about 10 bits, so use `Scale.F32`
12
- /// there. `scale=` also takes the name that `name` returns.
13
+ /// there. `scale=` also takes the name that `name` returns, and
14
+ /// `Scale(name)` gives the scale back, like `Scale("f16")`. Pickles load
15
+ /// through it, so `torch.load` accepts them once
16
+ /// `torch.serialization.add_safe_globals([Scale])` allows the class.
13
17
  #[pyclass(
14
18
  eq,
15
19
  frozen,
@@ -65,6 +69,11 @@ impl FromPyObject<'_, '_> for PyScale {
65
69
 
66
70
  #[pymethods]
67
71
  impl PyScale {
72
+ #[new]
73
+ fn new(name: PyScale) -> Self {
74
+ name
75
+ }
76
+
68
77
  /// The name that `scale=` also accepts: `'f32'`, `'f16'`, or `'bf16'`.
69
78
  #[getter]
70
79
  fn name(&self) -> &'static str {
@@ -79,13 +88,15 @@ impl PyScale {
79
88
  self.name()
80
89
  }
81
90
 
91
+ /// Pickles saved by quantize-py 0.3.0 and earlier call this with the name.
82
92
  #[staticmethod]
83
93
  fn _from_pickle(name: &str) -> PyResult<Self> {
84
94
  Self::from_name(name).ok_or_else(|| PyValueError::new_err("malformed pickle state"))
85
95
  }
86
96
 
87
- fn __reduce__<'py>(slf: &Bound<'py, Self>) -> PyResult<(Bound<'py, PyAny>, (&'static str,))> {
88
- let callable = slf.as_any().getattr("_from_pickle")?;
89
- Ok((callable, (slf.get().name(),)))
97
+ // Pickles call the class with the name, so that `torch.load` loads them
98
+ // once `add_safe_globals([Scale])` allows it, as with `Quantized`.
99
+ fn __reduce__<'py>(slf: &Bound<'py, Self>) -> (Bound<'py, PyType>, (&'static str,)) {
100
+ (slf.as_any().get_type(), (slf.get().name(),))
90
101
  }
91
102
  }
@@ -4,6 +4,7 @@ use pyo3::exceptions::PyValueError;
4
4
  use pyo3::prelude::*;
5
5
  use pyo3::types::PyType;
6
6
 
7
+ use crate::error::from_quantize;
7
8
  use crate::input::{bits_argument, block_argument};
8
9
  use crate::quantize_values;
9
10
  use crate::quantized::PyQuantized;
@@ -13,6 +14,17 @@ use crate::scale::PyScale;
13
14
  /// `Scheme.symmetric`, `Scheme.asymmetric`, and `Scheme.adaptive` build one,
14
15
  /// and `quantize` runs it. `Scheme.Q8_32` and `Scheme.Q4_32` are symmetric
15
16
  /// 8-bit and 4-bit codes, with blocks of 32.
17
+ ///
18
+ /// `Scheme(text)` reads a scheme written like a call to one of those
19
+ /// methods, without the `Scheme.`, such as
20
+ /// `Scheme("adaptive(block=32, tolerance=0.002)")`, or a constant's name,
21
+ /// such as `Scheme("Q4_32")`. Text that isn't a scheme, or that holds a
22
+ /// value `quantize` would reject, like `bits=99`, raises `QuantizeError`.
23
+ ///
24
+ /// Pickles and copies load through `Scheme(text)`, so a scheme that
25
+ /// `quantize` would reject, like `Scheme.symmetric(bits=99)`, raises the
26
+ /// same error when it's unpickled or copied. `torch.load` accepts pickles
27
+ /// once `torch.serialization.add_safe_globals([Scheme])` allows the class.
16
28
  #[pyclass(frozen, name = "Scheme", module = "quantize", eq, skip_from_py_object)]
17
29
  #[derive(Clone, Copy, PartialEq)]
18
30
  pub struct PyScheme {
@@ -35,6 +47,12 @@ impl PyScheme {
35
47
 
36
48
  #[pymethods]
37
49
  impl PyScheme {
50
+ #[new]
51
+ fn new(text: &str) -> PyResult<Self> {
52
+ let inner = text.parse().map_err(from_quantize)?;
53
+ Ok(Self { inner })
54
+ }
55
+
38
56
  #[classattr]
39
57
  #[pyo3(name = "Q8_32")]
40
58
  fn q8_32() -> Self {
@@ -165,6 +183,8 @@ impl PyScheme {
165
183
  self.pickle_parts()
166
184
  }
167
185
 
186
+ /// Pickles saved by quantize-py 0.3.0 and earlier call this with the
187
+ /// parts that `__getstate__` returns.
168
188
  #[staticmethod]
169
189
  fn _from_pickle(
170
190
  kind: &str,
@@ -186,8 +206,10 @@ impl PyScheme {
186
206
  }
187
207
  }
188
208
 
189
- fn __reduce__<'py>(slf: &Bound<'py, Self>) -> PyResult<(Bound<'py, PyAny>, SchemePickle)> {
190
- let callable = slf.getattr("_from_pickle")?;
191
- Ok((callable, slf.get().pickle_parts()))
209
+ // Pickles call the class with the scheme's text, so that `torch.load`
210
+ // loads them once `add_safe_globals([Scheme])` allows it, as with
211
+ // `Quantized`.
212
+ fn __reduce__<'py>(slf: &Bound<'py, Self>) -> (Bound<'py, PyType>, (String,)) {
213
+ (slf.get_type(), (slf.get().inner.to_string(),))
192
214
  }
193
215
  }
@@ -127,6 +127,16 @@ def test_except_quantize_error_catches_length():
127
127
  assert isinstance(raised.value, ValueError)
128
128
 
129
129
 
130
+ def test_a_copy_keeps_the_original_through_refine_and_alternate():
131
+ weights = np.linspace(-0.5, 0.5, 64, dtype=np.float32).reshape(2, 32)
132
+ for refit in [learned.refine, learned.alternate]:
133
+ quantized = quantize(weights, bits=4)
134
+ original = quantized.copy()
135
+ refit(quantized, weights)
136
+ assert quantized != original
137
+ assert original == quantize(weights, bits=4)
138
+
139
+
130
140
  def test_refine_and_alternate_work_while_another_thread_uses_the_tensor():
131
141
  weights = np.random.default_rng(0).standard_normal((256, 256)).astype(np.float32)
132
142
  quantized = quantize(weights, bits=4)
@@ -410,6 +410,19 @@ def test_scheme_factory_does_not_validate():
410
410
  assert scheme.bits == 1
411
411
  with pytest.raises(InvalidBitsError):
412
412
  scheme.quantize([0.1])
413
+ data = pickle.dumps(scheme)
414
+ with pytest.raises(InvalidBitsError):
415
+ pickle.loads(data)
416
+
417
+
418
+ def test_scheme_reads_text_written_like_its_class_methods():
419
+ assert Scheme("symmetric(bits=4, block=32)") == Scheme("Q4_32") == Scheme.Q4_32
420
+ assert Scheme("asymmetric(bits=3, block=7)") == Scheme.asymmetric(bits=3, block=7)
421
+ assert Scheme("adaptive(block=32, tolerance=0.002)") == Scheme.adaptive(tolerance=0.002)
422
+ with pytest.raises(QuantizeError, match=r"isn't a scheme; write one like symmetric\(bits=4"):
423
+ Scheme("symmetric:4:32")
424
+ with pytest.raises(InvalidBitsError):
425
+ Scheme("symmetric(bits=99, block=32)")
413
426
 
414
427
 
415
428
  def test_quantize_rejects_other_dimensions():
@@ -437,9 +450,12 @@ def test_scale_enum_selects_storage():
437
450
  def test_scale_can_be_given_by_name():
438
451
  for scale in [Scale.F32, Scale.F16, Scale.BF16]:
439
452
  assert quantize([0.1], scale=scale.name).scale == scale
453
+ assert Scale(scale.name) == scale
440
454
  assert Scale.BF16.name == "bf16"
441
455
  with pytest.raises(QuantizeError, match="scale must be a Scale or its name"):
442
456
  quantize([0.1], scale="float32")
457
+ with pytest.raises(QuantizeError, match="scale must be a Scale or its name"):
458
+ Scale("float32")
443
459
 
444
460
 
445
461
  def test_bad_values_raise_value_errors_that_name_the_argument():
@@ -521,25 +537,38 @@ def test_pickles_call_the_class_with_the_bytes_that_to_bytes_saves():
521
537
  assert Quantized(data) == Quantized.from_bytes(data) == quantized
522
538
 
523
539
 
524
- class OnlyQuantizedUnpickler(pickle.Unpickler):
540
+ class OnlyOurClassesUnpickler(pickle.Unpickler):
525
541
  """Stands in for `torch.load`, which by default refuses any global it
526
- doesn't trust, after `torch.serialization.add_safe_globals([Quantized])`.
542
+ doesn't trust, after
543
+ `torch.serialization.add_safe_globals([Quantized, Scale, Scheme])`.
527
544
  It also trusts `_codecs.encode`, as `torch.load` does, since pickles below
528
545
  protocol 3 store bytes with it."""
529
546
 
530
547
  def find_class(self, module, name):
531
- if (module, name) == ("quantize", "Quantized"):
532
- return Quantized
548
+ classes = {"Quantized": Quantized, "Scale": Scale, "Scheme": Scheme}
549
+ if module == "quantize" and name in classes:
550
+ return classes[name]
533
551
  if (module, name) == ("_codecs", "encode"):
534
552
  return codecs.encode
535
553
  raise pickle.UnpicklingError(f"{module}.{name} isn't allowed")
536
554
 
537
555
 
538
- def test_pickles_load_when_only_the_class_is_allowed_as_in_torch_load():
539
- checkpoint = {"layer": quantize(weight_matrix(8, 32), bits=4, scale=Scale.F16)}
556
+ def test_pickles_load_when_only_the_classes_are_allowed_as_in_torch_load():
557
+ checkpoint = {
558
+ "layer": quantize(weight_matrix(8, 32), bits=4, scale=Scale.F16),
559
+ "scheme": Scheme.adaptive(block=32, tolerance=0.002),
560
+ "scale": Scale.BF16,
561
+ }
540
562
  for protocol in range(pickle.HIGHEST_PROTOCOL + 1):
541
563
  data = pickle.dumps(checkpoint, protocol)
542
- assert OnlyQuantizedUnpickler(io.BytesIO(data)).load() == checkpoint
564
+ assert OnlyOurClassesUnpickler(io.BytesIO(data)).load() == checkpoint
565
+
566
+
567
+ def test_scales_and_schemes_pickle_as_a_call_to_their_class():
568
+ for value in [Scale.F16, Scheme.Q4_32, Scheme.adaptive(block=32, tolerance=0.002)]:
569
+ rebuild, arguments = value.__reduce__()
570
+ assert rebuild is type(value)
571
+ assert rebuild(*arguments) == value
543
572
 
544
573
 
545
574
  def test_from_bytes_rejects_bytes_that_do_not_hold_a_tensor():
@@ -599,6 +628,26 @@ def test_a_pickle_through_from_bytes_still_loads():
599
628
  assert pickle.loads(PICKLED_THROUGH_FROM_BYTES) == quantize(weights, bits=8, block=4)
600
629
 
601
630
 
631
+ # [Scale.F16, Scheme.adaptive(block=32, tolerance=0.002), Scheme.Q4_32], pickled
632
+ # by quantize-py 0.3.0, which called their _from_pickle methods.
633
+ SCALE_AND_SCHEMES_PICKLED_BY_0_3_0 = (
634
+ b"\x80\x04\x95\xad\x00\x00\x00\x00\x00\x00\x00]\x94(\x8c\x08builtins\x94\x8c"
635
+ b"\x07getattr\x94\x93\x94\x8c\x08quantize\x94\x8c\x05Scale\x94\x93\x94\x8c\x0c"
636
+ b"_from_pickle\x94\x86\x94R\x94\x8c\x03f16\x94\x85\x94R\x94h\x03\x8c\x08quanti"
637
+ b"ze\x94\x8c\x06Scheme\x94\x93\x94\x8c\x0c_from_pickle\x94\x86\x94R\x94(\x8c"
638
+ b"\x08adaptive\x94NK G?`bM\xe0\x00\x00\x00t\x94R\x94h\x12(\x8c\tsymmetric\x94K"
639
+ b"\x04K Nt\x94R\x94e."
640
+ )
641
+
642
+
643
+ def test_scales_and_schemes_pickled_by_0_3_0_still_load():
644
+ assert pickle.loads(SCALE_AND_SCHEMES_PICKLED_BY_0_3_0) == [
645
+ Scale.F16,
646
+ Scheme.adaptive(block=32, tolerance=0.002),
647
+ Scheme.Q4_32,
648
+ ]
649
+
650
+
602
651
  def test_quantized_compares_by_value():
603
652
  weights = weight_matrix(4, 32)
604
653
  quantized = quantize(weights, bits=4)