quantize-py 0.3.2__tar.gz → 0.4.0__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 (47) hide show
  1. {quantize_py-0.3.2 → quantize_py-0.4.0}/Cargo.lock +23 -2
  2. {quantize_py-0.3.2 → quantize_py-0.4.0}/Cargo.toml +4 -5
  3. {quantize_py-0.3.2 → quantize_py-0.4.0}/PKG-INFO +1 -1
  4. {quantize_py-0.3.2 → quantize_py-0.4.0}/README.md +2 -0
  5. {quantize_py-0.3.2 → quantize_py-0.4.0}/python/src/input.rs +21 -0
  6. {quantize_py-0.3.2 → quantize_py-0.4.0}/python/src/quantized/methods.rs +12 -4
  7. {quantize_py-0.3.2 → quantize_py-0.4.0}/python/src/quantized/parts.rs +18 -9
  8. {quantize_py-0.3.2 → quantize_py-0.4.0}/python/tests/test_quantize.py +29 -0
  9. {quantize_py-0.3.2 → quantize_py-0.4.0}/src/lib.rs +25 -6
  10. {quantize_py-0.3.2 → quantize_py-0.4.0}/src/methods/adaptive/mod.rs +8 -2
  11. {quantize_py-0.3.2 → quantize_py-0.4.0}/src/methods/asymmetric/mod.rs +17 -3
  12. {quantize_py-0.3.2 → quantize_py-0.4.0}/src/methods/symmetric/mod.rs +26 -10
  13. {quantize_py-0.3.2 → quantize_py-0.4.0}/src/shared/bytes.rs +19 -5
  14. quantize_py-0.4.0/src/shared/cores.rs +130 -0
  15. {quantize_py-0.3.2 → quantize_py-0.4.0}/src/shared/decode.rs +32 -57
  16. {quantize_py-0.3.2 → quantize_py-0.4.0}/src/shared/error.rs +29 -0
  17. {quantize_py-0.3.2 → quantize_py-0.4.0}/src/shared/mod.rs +2 -0
  18. {quantize_py-0.3.2 → quantize_py-0.4.0}/src/shared/scale.rs +5 -1
  19. {quantize_py-0.3.2 → quantize_py-0.4.0}/src/shared/tensor.rs +380 -38
  20. {quantize_py-0.3.2 → quantize_py-0.4.0}/LICENSE +0 -0
  21. {quantize_py-0.3.2 → quantize_py-0.4.0}/pyproject.toml +0 -0
  22. {quantize_py-0.3.2 → quantize_py-0.4.0}/python/Cargo.toml +0 -0
  23. {quantize_py-0.3.2 → quantize_py-0.4.0}/python/README.md +0 -0
  24. {quantize_py-0.3.2 → quantize_py-0.4.0}/python/quantize/__init__.py +0 -0
  25. {quantize_py-0.3.2 → quantize_py-0.4.0}/python/quantize/adaptive.py +0 -0
  26. {quantize_py-0.3.2 → quantize_py-0.4.0}/python/quantize/asymmetric.py +0 -0
  27. {quantize_py-0.3.2 → quantize_py-0.4.0}/python/quantize/learned.py +0 -0
  28. {quantize_py-0.3.2 → quantize_py-0.4.0}/python/quantize/symmetric.py +0 -0
  29. {quantize_py-0.3.2 → quantize_py-0.4.0}/python/src/error.rs +0 -0
  30. {quantize_py-0.3.2 → quantize_py-0.4.0}/python/src/gil.rs +0 -0
  31. {quantize_py-0.3.2 → quantize_py-0.4.0}/python/src/learned.rs +0 -0
  32. {quantize_py-0.3.2 → quantize_py-0.4.0}/python/src/lib.rs +0 -0
  33. {quantize_py-0.3.2 → quantize_py-0.4.0}/python/src/quantized/inner.rs +0 -0
  34. {quantize_py-0.3.2 → quantize_py-0.4.0}/python/src/quantized.rs +0 -0
  35. {quantize_py-0.3.2 → quantize_py-0.4.0}/python/src/scale.rs +0 -0
  36. {quantize_py-0.3.2 → quantize_py-0.4.0}/python/src/scheme.rs +0 -0
  37. {quantize_py-0.3.2 → quantize_py-0.4.0}/python/tests/test_learned.py +0 -0
  38. {quantize_py-0.3.2 → quantize_py-0.4.0}/src/kernels/block.rs +0 -0
  39. {quantize_py-0.3.2 → quantize_py-0.4.0}/src/kernels/i4.rs +0 -0
  40. {quantize_py-0.3.2 → quantize_py-0.4.0}/src/kernels/i8.rs +0 -0
  41. {quantize_py-0.3.2 → quantize_py-0.4.0}/src/kernels/mod.rs +0 -0
  42. {quantize_py-0.3.2 → quantize_py-0.4.0}/src/kernels/reduce.rs +0 -0
  43. {quantize_py-0.3.2 → quantize_py-0.4.0}/src/methods/learned/mod.rs +0 -0
  44. {quantize_py-0.3.2 → quantize_py-0.4.0}/src/methods/mod.rs +0 -0
  45. {quantize_py-0.3.2 → quantize_py-0.4.0}/src/shared/packed.rs +0 -0
  46. {quantize_py-0.3.2 → quantize_py-0.4.0}/src/shared/params.rs +0 -0
  47. {quantize_py-0.3.2 → quantize_py-0.4.0}/src/shared/scheme.rs +0 -0
@@ -1411,14 +1411,24 @@ dependencies = [
1411
1411
 
1412
1412
  [[package]]
1413
1413
  name = "quantize"
1414
- version = "0.3.2"
1414
+ version = "0.4.0"
1415
+ dependencies = [
1416
+ "half",
1417
+ "rayon",
1418
+ ]
1419
+
1420
+ [[package]]
1421
+ name = "quantize-files"
1422
+ version = "0.4.0"
1415
1423
  dependencies = [
1416
1424
  "half",
1425
+ "quantize",
1426
+ "serde_json",
1417
1427
  ]
1418
1428
 
1419
1429
  [[package]]
1420
1430
  name = "quantize-py"
1421
- version = "0.3.2"
1431
+ version = "0.4.0"
1422
1432
  dependencies = [
1423
1433
  "half",
1424
1434
  "numpy",
@@ -1778,6 +1788,17 @@ version = "1.16.2"
1778
1788
  source = "registry+https://github.com/rust-lang/crates.io-index"
1779
1789
  checksum = "f9395f0f0eee849a9b707b2f06bb92a6a422090e2123bb2ef8e87a0e61892a8e"
1780
1790
 
1791
+ [[package]]
1792
+ name = "smollm"
1793
+ version = "0.0.0"
1794
+ dependencies = [
1795
+ "hf-hub",
1796
+ "quantize",
1797
+ "quantize-files",
1798
+ "serde_json",
1799
+ "tokenizers",
1800
+ ]
1801
+
1781
1802
  [[package]]
1782
1803
  name = "socks"
1783
1804
  version = "0.3.4"
@@ -23,7 +23,7 @@ resolver = "3"
23
23
 
24
24
  [workspace.package]
25
25
  # Shared with the Python package. Between releases it ends in -dev.
26
- version = "0.3.2"
26
+ version = "0.4.0"
27
27
 
28
28
  [workspace.dependencies]
29
29
  half = "2"
@@ -34,13 +34,12 @@ numpy = "0.29"
34
34
  all = { level = "deny", priority = -1 }
35
35
 
36
36
  [features]
37
- default = ["std"]
38
- # Does nothing, since the crate always needs the standard library. Kept so
39
- # that a Cargo.toml that names it still builds.
40
- std = []
37
+ # Splits large matmul calls across every core.
38
+ rayon = ["dep:rayon"]
41
39
 
42
40
  [dependencies]
43
41
  half = { workspace = true }
42
+ rayon = { version = "1", optional = true }
44
43
 
45
44
  [lints]
46
45
  workspace = true
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: quantize-py
3
- Version: 0.3.2
3
+ Version: 0.4.0
4
4
  Classifier: Programming Language :: Python :: 3
5
5
  Classifier: Programming Language :: Python :: 3.12
6
6
  Classifier: Programming Language :: Python :: 3.13
@@ -48,6 +48,8 @@ let dot = q.dot(&weights).unwrap(); // 0.926
48
48
 
49
49
  for a weight matrix, `q.set_shape(rows, columns)` records its shape, so `q.matmul(&inputs)` can multiply a batch of inputs by it, like a linear layer, and `q.dequantize_row_into(row, &mut out)` can decode one row, like an embedding lookup. `q.to_bytes()` saves a tensor, shape included, and `Quantized::from_bytes` loads it back.
50
50
 
51
+ to run `matmul` on every core, turn on the `rayon` feature with `cargo add quantize --features rayon`. a big enough call splits across the cores, by the matrix's rows, by vectors, or both, and the results are the same, bit for bit. smaller calls stay on one core. a program that already keeps every core busy calling `matmul` loses up to about a quarter, most of all for one vector through a large matrix, and one with cores to spare gains. without it, the only dependency is `half`.
52
+
51
53
  codes can be 2 to 16 bits, and scales `f32`, `f16`, or `bf16`. the rest is in the [api docs](https://docs.rs/quantize). for python, see [python/](https://github.com/aksheyd/quantize/tree/main/python).
52
54
 
53
55
  ---
@@ -195,6 +195,27 @@ pub fn as_packed_codes(obj: &Bound<'_, PyAny>) -> PyResult<Vec<u8>> {
195
195
  read_uint8(obj, PACKED_CODES_TYPE)
196
196
  }
197
197
 
198
+ /// Read block widths: a 1-D uint8 array, like `Quantized.block_bits`
199
+ /// returns, or a sequence of ints. The array is copied at once, since reading
200
+ /// it as a sequence would convert one NumPy scalar at a time.
201
+ pub fn as_block_bits(obj: &Bound<'_, PyAny>) -> PyResult<Vec<u32>> {
202
+ match obj.cast::<PyArray1<u8>>() {
203
+ Ok(widths) => {
204
+ let widths = widths.try_readonly()?;
205
+ Ok(widths
206
+ .as_array()
207
+ .iter()
208
+ .map(|&bits| u32::from(bits))
209
+ .collect())
210
+ }
211
+ Err(_) => obj.extract().inspect_err(|error: &PyErr| {
212
+ // Errors like "Can't extract `str` to `Vec`" don't name the
213
+ // argument, so add the note PyO3 adds to the arguments it reads.
214
+ let _ = error.add_note(obj.py(), "while processing 'block_bits'");
215
+ }),
216
+ }
217
+ }
218
+
198
219
  /// Read the bytes that `to_bytes` saved: `bytes` or another bytes-like
199
220
  /// object, or a 1-D uint8 array.
200
221
  pub fn as_bytes(obj: &Bound<'_, PyAny>) -> PyResult<Vec<u8>> {
@@ -10,7 +10,7 @@ use super::parts::Parts;
10
10
  use crate::error::from_quantize;
11
11
  use crate::gil::detach_if_large;
12
12
  use crate::input::{
13
- as_bytes, as_f32_array, as_f32_matmul_values, as_f32_values, as_packed_codes,
13
+ as_block_bits, as_bytes, as_f32_array, as_f32_matmul_values, as_f32_values, as_packed_codes,
14
14
  as_writable_f32_out, check_shape,
15
15
  };
16
16
  use crate::scale::PyScale;
@@ -235,14 +235,18 @@ impl PyQuantized {
235
235
  .map(|bits| bits.to_vec().into_pyarray(py)))
236
236
  }
237
237
 
238
- /// Bytes held by the codes, scales, zero-points, and block widths.
238
+ /// Bytes held by the codes, scales, zero-points, and block widths. An
239
+ /// adaptive matrix also keeps 8 bytes a row in memory, to find where each
240
+ /// row starts, which this doesn't count.
239
241
  #[getter]
240
242
  fn nbytes(&self) -> usize {
241
243
  with_inner!(&self.snapshot(), |quantized| quantized.nbytes())
242
244
  }
243
245
 
244
246
  /// Bits per value, counting the scales: 4-bit codes with one f16 scale
245
- /// per 32 values cost 4.5.
247
+ /// per 32 values cost 4.5. An adaptive matrix also keeps 8 bytes a row in
248
+ /// memory, to find where each row starts, which this doesn't count: at
249
+ /// 64 columns, that's 1 more bit per value.
246
250
  #[getter]
247
251
  fn bits_per_element(&self) -> f32 {
248
252
  with_inner!(&self.snapshot(), |quantized| quantized.bits_per_element())
@@ -270,8 +274,12 @@ impl PyQuantized {
270
274
  scale: PyScale,
271
275
  zero_points: Option<Bound<'_, PyAny>>,
272
276
  bits: Option<u32>,
273
- block_bits: Option<Vec<u32>>,
277
+ block_bits: Option<Bound<'_, PyAny>>,
274
278
  ) -> PyResult<Self> {
279
+ let block_bits = match block_bits {
280
+ Some(block_bits) => Some(as_block_bits(&block_bits)?),
281
+ None => None,
282
+ };
275
283
  let zero_points = match zero_points {
276
284
  Some(zero_points) => as_f32_values(&zero_points)?.to_vec()?,
277
285
  None => Vec::new(),
@@ -23,14 +23,19 @@ pub(crate) struct Parts {
23
23
  }
24
24
 
25
25
  impl Parts {
26
- /// Build the variant that `kind` names, then [`Quantized::validate`] it,
27
- /// so parts that don't fit together are an error instead of a tensor
28
- /// that decodes out of bounds.
26
+ /// Build the variant that `kind` names as a flat vector, then
27
+ /// [`Quantized::validate`] it, so parts that don't fit together are an
28
+ /// error instead of a tensor that decodes out of bounds. A 2-D `shape`
29
+ /// goes on last, through `set_shape`, which finds where each adaptive row
30
+ /// starts.
29
31
  pub(crate) fn into_quantized<S: Scale>(self) -> PyResult<Quantized<S>> {
30
- let (len, columns) = match self.shape[..] {
32
+ let (len, rows_and_columns) = match self.shape[..] {
31
33
  [len] => (len, None),
34
+ // Named here, since validating the empty tensor first would blame
35
+ // its scales instead.
36
+ [_, 0] => return Err(from_quantize(Error::ShapeMismatch { len: 0, columns: 0 })),
32
37
  [rows, columns] => match rows.checked_mul(columns) {
33
- Some(len) => (len, Some(columns)),
38
+ Some(len) => (len, Some((rows, columns))),
34
39
  None => return Err(PyErr::new::<QuantizeError, _>(SHAPE_TOO_LARGE)),
35
40
  },
36
41
  _ => return Err(PyErr::new::<QuantizeError, _>(SHAPE_DIMENSIONS)),
@@ -39,13 +44,13 @@ impl Parts {
39
44
  let scales: Vec<S> = self.scales.into_iter().map(S::from_f32).collect();
40
45
  let zero_points: Vec<S> = self.zero_points.into_iter().map(S::from_f32).collect();
41
46
 
42
- let quantized = match (self.kind.as_str(), self.bits, self.block_bits) {
47
+ let mut quantized = match (self.kind.as_str(), self.bits, self.block_bits) {
43
48
  ("symmetric", Some(bits), None) if zero_points.is_empty() => Quantized::Symmetric {
44
49
  scales,
45
50
  codes: packed_codes(self.codes, bits, len)?,
46
51
  block,
47
52
  len,
48
- columns,
53
+ columns: None,
49
54
  },
50
55
  ("asymmetric", Some(bits), None) => Quantized::Asymmetric {
51
56
  scales,
@@ -53,7 +58,7 @@ impl Parts {
53
58
  codes: packed_codes(self.codes, bits, len)?,
54
59
  block,
55
60
  len,
56
- columns,
61
+ columns: None,
57
62
  },
58
63
  ("adaptive", None, Some(block_bits)) => Quantized::Adaptive {
59
64
  scales,
@@ -62,11 +67,15 @@ impl Parts {
62
67
  block_bits: block_widths(block_bits)?,
63
68
  block,
64
69
  len,
65
- columns,
70
+ columns: None,
71
+ row_starts: Vec::new(),
66
72
  },
67
73
  _ => return Err(PyErr::new::<QuantizeError, _>(KIND_PARTS)),
68
74
  };
69
75
  quantized.validate().map_err(from_quantize)?;
76
+ if let Some((rows, columns)) = rows_and_columns {
77
+ quantized.set_shape(rows, columns).map_err(from_quantize)?;
78
+ }
70
79
  Ok(quantized)
71
80
  }
72
81
  }
@@ -759,6 +759,7 @@ def test_from_parts_rebuilds_parts_saved_with_numpy_as_the_readme_says():
759
759
  quantize(weights, bits=4, block=32, scale=scale),
760
760
  asymmetric.quantize(weights, bits=5, block=16, scale=scale),
761
761
  adaptive.quantize(weights.ravel(), block=32, tolerance=0.001, scale=scale),
762
+ adaptive.quantize(weights, block=32, tolerance=0.001, scale=scale),
762
763
  quantize([], scale=scale),
763
764
  ]:
764
765
  parts = {
@@ -826,3 +827,31 @@ def test_adaptive_block_widths_take_one_byte_each():
826
827
  block_bits=[300, 8, 8],
827
828
  scale="f32",
828
829
  )
830
+
831
+
832
+ def test_from_parts_reads_block_bits_from_an_array_or_any_sequence_of_ints():
833
+ quantized = adaptive.quantize(weight_matrix(3, 30), block=8, tolerance=0.001)
834
+ parts = {
835
+ "kind": "adaptive",
836
+ "shape": quantized.shape,
837
+ "block": quantized.block,
838
+ "codes": quantized.codes,
839
+ "scales": quantized.scales,
840
+ "zero_points": quantized.zero_points,
841
+ "scale": "f32",
842
+ }
843
+ widths = quantized.block_bits
844
+ non_contiguous_widths = np.repeat(widths, 2)[::2]
845
+ for block_bits in [
846
+ widths,
847
+ non_contiguous_widths,
848
+ widths.tolist(),
849
+ tuple(widths),
850
+ widths.astype(np.int64),
851
+ ]:
852
+ assert Quantized.from_parts(**parts, block_bits=block_bits) == quantized
853
+ with pytest.raises(InvalidBitsError, match="got 17"):
854
+ Quantized.from_parts(**parts, block_bits=np.full_like(widths, 17))
855
+ with pytest.raises(TypeError) as raised:
856
+ Quantized.from_parts(**parts, block_bits="4")
857
+ assert raised.value.__notes__ == ["while processing 'block_bits'"]
@@ -25,9 +25,10 @@
25
25
  //! 32 values adds 16 / 32 = 0.5 bits to each value: 4-bit codes cost 4.5 bits
26
26
  //! per value.
27
27
  //!
28
- //! `BITS` and `BLOCK` are const generics. To choose them at run time, call the
29
- //! scheme's `quantize_with`, like [`symmetric::quantize_with`], which takes
30
- //! them as ordinary arguments.
28
+ //! `BITS` and `BLOCK` are const generics, so a width outside 2 to 16 or a
29
+ //! block of 0 stops the build instead of returning an error. To choose them at
30
+ //! run time, call the scheme's `quantize_with`, like
31
+ //! [`symmetric::quantize_with`], which takes them as ordinary arguments.
31
32
  //!
32
33
  //! [`quantize`] is symmetric: each block gets one scale. Everything below
33
34
  //! uses the same [`Quantized`] type:
@@ -64,9 +65,25 @@
64
65
  //!
65
66
  //! ## Features
66
67
  //!
67
- //! The crate needs the standard library. Its `std` feature, on by default,
68
- //! does nothing, and is kept so that a `Cargo.toml` that names it still
69
- //! builds.
68
+ //! The `rayon` feature, off by default, splits a large
69
+ //! [`matmul`](Quantized::matmul) or [`matmul_into`](Quantized::matmul_into)
70
+ //! across every core with the [rayon](https://docs.rs/rayon) crate, and the
71
+ //! results are the same, bit for bit. A call splits the matrix's rows, its
72
+ //! inputs, or both, as `matmul` explains. A call too small to split stays on
73
+ //! the thread that made it, as every call does without the feature.
74
+ //!
75
+ //! A program that already keeps every core busy calling `matmul` has no core
76
+ //! free for a share, so a split can't make its calls finish sooner. Its
77
+ //! smaller calls run as they did, and a call big enough to split costs up to
78
+ //! about a quarter more, the most for one input through a large matrix. A
79
+ //! program with cores to spare gains instead, since the shares run on them.
80
+ //! Rayon starts one thread per core: set `RAYON_NUM_THREADS` to use fewer, or
81
+ //! to 1 to stop splitting, which holds for any other rayon code in the program
82
+ //! too. Cargo turns a feature on for every user of a crate once anything in
83
+ //! the build asks for it, so a dependency can turn this one on for you.
84
+ //!
85
+ //! Turn it on with `cargo add quantize --features rayon`. Without it, the
86
+ //! crate's only dependency is `half`.
70
87
 
71
88
  #![warn(missing_docs)]
72
89
 
@@ -87,4 +104,6 @@ pub use shared::scheme::Scheme;
87
104
  pub use shared::tensor::Quantized;
88
105
  pub use symmetric::{quantize, quantize_tensor};
89
106
 
107
+ #[cfg(feature = "rayon")]
108
+ pub(crate) use shared::cores;
90
109
  pub(crate) use shared::{decode, error, packed, scale, tensor};
@@ -24,21 +24,25 @@ use crate::tensor::Quantized;
24
24
  /// times past on blocks far from zero, as [`asymmetric`](crate::asymmetric)
25
25
  /// explains.
26
26
  ///
27
+ /// `BLOCK` must be at least 1, or the build stops, as in
28
+ /// [`symmetric::quantize`](crate::symmetric::quantize).
29
+ ///
27
30
  /// # Errors
28
31
  ///
29
32
  /// [`Error::InvalidTolerance`] if `tolerance` is not finite and `> 0`.
30
33
  /// [`Error::ToleranceTooTight`] if even 8 bits can't round a block within
31
34
  /// `tolerance`, with the smallest tolerance that every block meets.
32
- /// [`Error::InvalidBlock`] if `BLOCK == 0`.
33
35
  /// [`Error::ScaleOutOfRange`] if `S` can't hold a block's scale or zero-point.
34
36
  pub fn quantize<S: Scale, const BLOCK: usize>(
35
37
  values: &[f32],
36
38
  tolerance: f32,
37
39
  ) -> Result<Quantized<S>> {
40
+ const { assert!(BLOCK >= 1, "BLOCK must be at least 1") };
38
41
  quantize_with::<S>(values, BLOCK, tolerance)
39
42
  }
40
43
 
41
- /// Runtime-block variant of [`quantize`].
44
+ /// [`quantize`] with the block size chosen at run time. It returns the same
45
+ /// errors, plus [`Error::InvalidBlock`] if `block` is 0.
42
46
  pub fn quantize_with<S: Scale>(
43
47
  values: &[f32],
44
48
  block: usize,
@@ -55,6 +59,7 @@ pub fn quantize_with<S: Scale>(
55
59
  block,
56
60
  len: 0,
57
61
  columns: None,
62
+ row_starts: Vec::new(),
58
63
  });
59
64
  }
60
65
 
@@ -98,6 +103,7 @@ pub fn quantize_with<S: Scale>(
98
103
  block,
99
104
  len: values.len(),
100
105
  columns: None,
106
+ row_starts: Vec::new(),
101
107
  })
102
108
  }
103
109
 
@@ -15,18 +15,30 @@ use crate::tensor::Quantized;
15
15
 
16
16
  /// Quantize into blocks of `BLOCK` using one bit width for every block.
17
17
  ///
18
+ /// `BITS` must be from 2 to 16, and `BLOCK` at least 1, or the build stops,
19
+ /// as in [`symmetric::quantize`](crate::symmetric::quantize).
20
+ ///
18
21
  /// # Errors
19
22
  ///
20
- /// [`crate::Error::InvalidBits`] or [`crate::Error::InvalidBlock`], and
21
23
  /// [`crate::Error::ScaleOutOfRange`] if `S` can't hold a block's scale or
22
24
  /// zero-point.
23
25
  pub fn quantize<S: Scale, const BITS: u32, const BLOCK: usize>(
24
26
  values: &[f32],
25
27
  ) -> Result<Quantized<S>> {
28
+ const {
29
+ assert!(2 <= BITS && BITS <= 16, "BITS must be from 2 to 16");
30
+ assert!(BLOCK >= 1, "BLOCK must be at least 1");
31
+ }
26
32
  quantize_with::<S>(values, BITS, BLOCK)
27
33
  }
28
34
 
29
- /// Runtime-width variant of [`quantize`].
35
+ /// [`quantize`] with the bit width and block size chosen at run time.
36
+ ///
37
+ /// # Errors
38
+ ///
39
+ /// [`crate::Error::InvalidBits`] or [`crate::Error::InvalidBlock`], and
40
+ /// [`crate::Error::ScaleOutOfRange`] if `S` can't hold a block's scale or
41
+ /// zero-point.
30
42
  pub fn quantize_with<S: Scale>(values: &[f32], bits: u32, block: usize) -> Result<Quantized<S>> {
31
43
  check_bits(bits)?;
32
44
  check_block(block)?;
@@ -58,8 +70,10 @@ pub fn quantize_with<S: Scale>(values: &[f32], bits: u32, block: usize) -> Resul
58
70
  })
59
71
  }
60
72
 
61
- /// Quantize the entire tensor with one scale and one zero-point.
73
+ /// Quantize the entire tensor with one scale and one zero-point. `BITS` must
74
+ /// be from 2 to 16, as in [`quantize`].
62
75
  pub fn quantize_tensor<S: Scale, const BITS: u32>(values: &[f32]) -> Result<Quantized<S>> {
76
+ const { assert!(2 <= BITS && BITS <= 16, "BITS must be from 2 to 16") };
63
77
  quantize_with::<S>(values, BITS, values.len().max(1))
64
78
  }
65
79
 
@@ -8,17 +8,32 @@ use crate::tensor::Quantized;
8
8
 
9
9
  /// Quantize `values` into fixed-size blocks of `BLOCK` with `BITS`-wide codes.
10
10
  ///
11
+ /// `BITS` must be from 2 to 16, and `BLOCK` at least 1. Anything else stops
12
+ /// the build, though `cargo check` doesn't catch it:
13
+ ///
14
+ /// ```compile_fail
15
+ /// let q = quantize::quantize::<f32, 1, 32>(&[0.5]);
16
+ /// ```
17
+ ///
11
18
  /// # Errors
12
19
  ///
13
- /// [`crate::Error::InvalidBits`] or [`crate::Error::InvalidBlock`], and
14
20
  /// [`crate::Error::ScaleOutOfRange`] if `S` can't hold a block's scale.
15
21
  pub fn quantize<S: Scale, const BITS: u32, const BLOCK: usize>(
16
22
  values: &[f32],
17
23
  ) -> Result<Quantized<S>> {
24
+ const {
25
+ assert!(2 <= BITS && BITS <= 16, "BITS must be from 2 to 16");
26
+ assert!(BLOCK >= 1, "BLOCK must be at least 1");
27
+ }
18
28
  quantize_with::<S>(values, BITS, BLOCK)
19
29
  }
20
30
 
21
- /// Runtime-width variant of [`quantize`].
31
+ /// [`quantize`] with the bit width and block size chosen at run time.
32
+ ///
33
+ /// # Errors
34
+ ///
35
+ /// [`crate::Error::InvalidBits`] or [`crate::Error::InvalidBlock`], and
36
+ /// [`crate::Error::ScaleOutOfRange`] if `S` can't hold a block's scale.
22
37
  pub fn quantize_with<S: Scale>(values: &[f32], bits: u32, block: usize) -> Result<Quantized<S>> {
23
38
  check_bits(bits)?;
24
39
  check_block(block)?;
@@ -41,8 +56,10 @@ pub fn quantize_with<S: Scale>(values: &[f32], bits: u32, block: usize) -> Resul
41
56
  })
42
57
  }
43
58
 
44
- /// Quantize the entire tensor with a single scale.
59
+ /// Quantize the entire tensor with a single scale. `BITS` must be from 2 to
60
+ /// 16, as in [`quantize`].
45
61
  pub fn quantize_tensor<S: Scale, const BITS: u32>(values: &[f32]) -> Result<Quantized<S>> {
62
+ const { assert!(2 <= BITS && BITS <= 16, "BITS must be from 2 to 16") };
46
63
  quantize_with::<S>(values, BITS, values.len().max(1))
47
64
  }
48
65
 
@@ -379,13 +396,12 @@ mod tests {
379
396
  let mut q = quantize::<f32, 8, 32>(&w).unwrap();
380
397
  q.set_shape(2, 32).unwrap();
381
398
  let inputs = [0.0_f32; 48];
382
- assert!(matches!(
383
- q.matmul(&inputs),
384
- Err(crate::Error::ShapeMismatch {
385
- len: 48,
386
- columns: 32
387
- })
388
- ));
399
+ let mismatch = crate::Error::InputMismatch {
400
+ columns: 32,
401
+ got: 48,
402
+ };
403
+ assert_eq!(q.matmul(&inputs), Err(mismatch.clone()));
404
+ assert_eq!(q.matmul_into(&inputs, &mut [0.0; 2]), Err(mismatch));
389
405
  }
390
406
 
391
407
  #[test]
@@ -113,7 +113,10 @@ impl<S: Scale> Quantized<S> {
113
113
 
114
114
  check_block(block)?;
115
115
  let blocks = len.div_ceil(block);
116
- let quantized = match kind {
116
+ // Every kind loads as a flat vector first. The shape goes on last,
117
+ // through `set_shape`, since it finds where each adaptive row starts
118
+ // by reading widths that must be checked first.
119
+ let mut quantized = match kind {
117
120
  SYMMETRIC => {
118
121
  let scales = reader.scales(blocks)?;
119
122
  let codes = reader.codes(code_bits, len)?;
@@ -122,7 +125,7 @@ impl<S: Scale> Quantized<S> {
122
125
  codes,
123
126
  block,
124
127
  len,
125
- columns,
128
+ columns: None,
126
129
  }
127
130
  }
128
131
  ASYMMETRIC => {
@@ -135,7 +138,7 @@ impl<S: Scale> Quantized<S> {
135
138
  codes,
136
139
  block,
137
140
  len,
138
- columns,
141
+ columns: None,
139
142
  }
140
143
  }
141
144
  ADAPTIVE => {
@@ -150,7 +153,8 @@ impl<S: Scale> Quantized<S> {
150
153
  block_bits,
151
154
  block,
152
155
  len,
153
- columns,
156
+ columns: None,
157
+ row_starts: Vec::new(),
154
158
  }
155
159
  }
156
160
  _ => return Err(malformed("unknown kind")),
@@ -164,6 +168,12 @@ impl<S: Scale> Quantized<S> {
164
168
  Ordering::Equal => {}
165
169
  }
166
170
  quantized.validate()?;
171
+ if let Some(columns) = columns {
172
+ if !len.is_multiple_of(columns) {
173
+ return Err(Error::ShapeMismatch { len, columns });
174
+ }
175
+ quantized.set_shape(len / columns, columns)?;
176
+ }
167
177
  Ok(quantized)
168
178
  }
169
179
  }
@@ -233,11 +243,15 @@ mod tests {
233
243
  five_by_sixteen.set_shape(5, 16).unwrap();
234
244
  let mut two_by_forty = asymmetric::quantize_with(&values, 8, 32).unwrap();
235
245
  two_by_forty.set_shape(2, 40).unwrap();
236
- let tensors: [Quantized<S>; 5] = [
246
+ // Its second row starts partway through a block.
247
+ let mut adaptive_two_by_forty = adaptive::quantize_with(&values, 32, 0.01).unwrap();
248
+ adaptive_two_by_forty.set_shape(2, 40).unwrap();
249
+ let tensors: [Quantized<S>; 6] = [
237
250
  symmetric::quantize_with(&values, 4, 32).unwrap(),
238
251
  five_by_sixteen,
239
252
  two_by_forty,
240
253
  adaptive::quantize_with(&values, 32, 0.01).unwrap(),
254
+ adaptive_two_by_forty,
241
255
  symmetric::quantize_with(&[], 8, 32).unwrap(),
242
256
  ];
243
257
  for quantized in tensors {