pyturboquant-cpu 0.1.0__py3-none-any.whl
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.
- pyturboquant_cpu/__init__.py +36 -0
- pyturboquant_cpu/codebook.py +182 -0
- pyturboquant_cpu/packing.py +258 -0
- pyturboquant_cpu/qjl.py +86 -0
- pyturboquant_cpu/quantizer.py +415 -0
- pyturboquant_cpu/rotation.py +84 -0
- pyturboquant_cpu-0.1.0.dist-info/METADATA +170 -0
- pyturboquant_cpu-0.1.0.dist-info/RECORD +11 -0
- pyturboquant_cpu-0.1.0.dist-info/WHEEL +5 -0
- pyturboquant_cpu-0.1.0.dist-info/licenses/LICENSE +176 -0
- pyturboquant_cpu-0.1.0.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,36 @@
|
|
|
1
|
+
"""
|
|
2
|
+
pyturboquant-cpu: CPU implementation of TurboQuant vector quantization.
|
|
3
|
+
|
|
4
|
+
TurboQuant is a data-oblivious vector quantization algorithm from
|
|
5
|
+
"TurboQuant: Online Vector Quantization with Near-optimal Distortion Rate"
|
|
6
|
+
(Zandieh et al., arXiv:2504.19874).
|
|
7
|
+
|
|
8
|
+
Two quantization modes are provided:
|
|
9
|
+
|
|
10
|
+
- **MSE-optimal** (``quantize_mse`` / ``dequantize_mse``):
|
|
11
|
+
Minimises mean-squared reconstruction error.
|
|
12
|
+
- **Inner-product-optimal** (``quantize_prod`` / ``dequantize_prod``):
|
|
13
|
+
Gives an *unbiased* inner-product estimator by combining an MSE
|
|
14
|
+
quantizer with a 1-bit QJL residual correction.
|
|
15
|
+
"""
|
|
16
|
+
|
|
17
|
+
from pyturboquant_cpu.quantizer import (
|
|
18
|
+
QuantizedMSE,
|
|
19
|
+
QuantizedProd,
|
|
20
|
+
quantize_mse,
|
|
21
|
+
dequantize_mse,
|
|
22
|
+
quantize_prod,
|
|
23
|
+
dequantize_prod,
|
|
24
|
+
)
|
|
25
|
+
|
|
26
|
+
__version__ = "0.1.0"
|
|
27
|
+
|
|
28
|
+
__all__ = [
|
|
29
|
+
"quantize_mse",
|
|
30
|
+
"dequantize_mse",
|
|
31
|
+
"quantize_prod",
|
|
32
|
+
"dequantize_prod",
|
|
33
|
+
"QuantizedMSE",
|
|
34
|
+
"QuantizedProd",
|
|
35
|
+
"__version__",
|
|
36
|
+
]
|
|
@@ -0,0 +1,182 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Lloyd-Max codebook computation for the Beta coordinate distribution.
|
|
3
|
+
|
|
4
|
+
After a random orthogonal rotation, each coordinate of a unit-norm vector
|
|
5
|
+
in R^d follows a Beta((d-1)/2, (d-1)/2) distribution on [-1, 1].
|
|
6
|
+
|
|
7
|
+
This module solves the 1-D Lloyd-Max (k-means) quantisation problem for
|
|
8
|
+
that distribution and caches the results for reuse.
|
|
9
|
+
"""
|
|
10
|
+
|
|
11
|
+
from __future__ import annotations
|
|
12
|
+
|
|
13
|
+
import functools
|
|
14
|
+
import math
|
|
15
|
+
from typing import Tuple
|
|
16
|
+
|
|
17
|
+
import numpy as np
|
|
18
|
+
from scipy import integrate, special
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
def beta_pdf(x: np.ndarray | float, dim: int) -> np.ndarray | float:
|
|
22
|
+
"""Probability density of a single coordinate after random rotation.
|
|
23
|
+
|
|
24
|
+
For a unit vector uniformly distributed on S^{d-1}, each coordinate
|
|
25
|
+
follows the distribution:
|
|
26
|
+
|
|
27
|
+
f(x) = Γ(d/2) / (√π · Γ((d-1)/2)) · (1 - x²)^{(d-3)/2}
|
|
28
|
+
|
|
29
|
+
for x ∈ [-1, 1].
|
|
30
|
+
|
|
31
|
+
Parameters
|
|
32
|
+
----------
|
|
33
|
+
x : array_like
|
|
34
|
+
Points at which to evaluate the PDF.
|
|
35
|
+
dim : int
|
|
36
|
+
Ambient dimension *d*.
|
|
37
|
+
|
|
38
|
+
Returns
|
|
39
|
+
-------
|
|
40
|
+
array_like
|
|
41
|
+
PDF values.
|
|
42
|
+
"""
|
|
43
|
+
x = np.asarray(x, dtype=np.float64)
|
|
44
|
+
d = dim
|
|
45
|
+
|
|
46
|
+
# Use log-gamma for numerical stability
|
|
47
|
+
log_norm = (
|
|
48
|
+
special.gammaln(d / 2.0)
|
|
49
|
+
- 0.5 * math.log(math.pi)
|
|
50
|
+
- special.gammaln((d - 1) / 2.0)
|
|
51
|
+
)
|
|
52
|
+
|
|
53
|
+
exponent = (d - 3) / 2.0
|
|
54
|
+
|
|
55
|
+
# Mask values outside support
|
|
56
|
+
mask = np.abs(x) < 1.0
|
|
57
|
+
result = np.zeros_like(x, dtype=np.float64)
|
|
58
|
+
if exponent == 0:
|
|
59
|
+
# d = 3: uniform on [-1, 1]
|
|
60
|
+
result[mask] = math.exp(log_norm)
|
|
61
|
+
else:
|
|
62
|
+
safe_x = np.where(mask, x, 0.0)
|
|
63
|
+
log_body = exponent * np.log1p(-safe_x**2)
|
|
64
|
+
result[mask] = np.exp(log_norm + log_body[mask])
|
|
65
|
+
|
|
66
|
+
return result
|
|
67
|
+
|
|
68
|
+
|
|
69
|
+
def _conditional_expectation(a: float, b: float, dim: int) -> float:
|
|
70
|
+
"""Compute E[X | a ≤ X ≤ b] for the Beta coordinate distribution.
|
|
71
|
+
|
|
72
|
+
Parameters
|
|
73
|
+
----------
|
|
74
|
+
a, b : float
|
|
75
|
+
Interval endpoints within [-1, 1].
|
|
76
|
+
dim : int
|
|
77
|
+
Ambient dimension *d*.
|
|
78
|
+
|
|
79
|
+
Returns
|
|
80
|
+
-------
|
|
81
|
+
float
|
|
82
|
+
Conditional mean.
|
|
83
|
+
"""
|
|
84
|
+
if b - a < 1e-15:
|
|
85
|
+
return (a + b) / 2.0
|
|
86
|
+
|
|
87
|
+
# Numerator: ∫_a^b x · f(x) dx
|
|
88
|
+
numerator, _ = integrate.quad(lambda x: x * beta_pdf(x, dim), a, b)
|
|
89
|
+
# Denominator: ∫_a^b f(x) dx
|
|
90
|
+
denominator, _ = integrate.quad(lambda x: beta_pdf(x, dim), a, b)
|
|
91
|
+
|
|
92
|
+
if abs(denominator) < 1e-30:
|
|
93
|
+
return (a + b) / 2.0
|
|
94
|
+
|
|
95
|
+
return numerator / denominator
|
|
96
|
+
|
|
97
|
+
|
|
98
|
+
def lloyd_max_codebook(
|
|
99
|
+
dim: int,
|
|
100
|
+
bits: int,
|
|
101
|
+
max_iter: int = 300,
|
|
102
|
+
tol: float = 1e-14,
|
|
103
|
+
) -> Tuple[np.ndarray, np.ndarray]:
|
|
104
|
+
"""Compute Lloyd-Max optimal codebook for the Beta coordinate distribution.
|
|
105
|
+
|
|
106
|
+
Solves the 1-D k-means problem for the distribution induced by random
|
|
107
|
+
rotation on the unit sphere in dimension *dim*, using *2^bits* centroids.
|
|
108
|
+
|
|
109
|
+
Parameters
|
|
110
|
+
----------
|
|
111
|
+
dim : int
|
|
112
|
+
Vector dimension *d*.
|
|
113
|
+
bits : int
|
|
114
|
+
Bit-width per coordinate. Produces ``2**bits`` centroids.
|
|
115
|
+
max_iter : int, optional
|
|
116
|
+
Maximum Lloyd iterations.
|
|
117
|
+
tol : float, optional
|
|
118
|
+
Convergence tolerance on centroid movement.
|
|
119
|
+
|
|
120
|
+
Returns
|
|
121
|
+
-------
|
|
122
|
+
centroids : ndarray, shape (2**bits,)
|
|
123
|
+
Sorted codebook centroids.
|
|
124
|
+
boundaries : ndarray, shape (2**bits - 1,)
|
|
125
|
+
Decision boundaries (midpoints between consecutive centroids).
|
|
126
|
+
"""
|
|
127
|
+
n_levels = 1 << bits # 2^bits
|
|
128
|
+
|
|
129
|
+
# Initialise centroids uniformly in [-1, 1]
|
|
130
|
+
centroids = np.linspace(-1.0, 1.0, n_levels + 2)[1:-1].copy()
|
|
131
|
+
|
|
132
|
+
for _ in range(max_iter):
|
|
133
|
+
# Boundaries = midpoints between consecutive centroids
|
|
134
|
+
boundaries = 0.5 * (centroids[:-1] + centroids[1:])
|
|
135
|
+
|
|
136
|
+
# Build full interval edges: [-1, b_1, b_2, ..., b_{n-1}, 1]
|
|
137
|
+
edges = np.empty(n_levels + 1, dtype=np.float64)
|
|
138
|
+
edges[0] = -1.0
|
|
139
|
+
edges[-1] = 1.0
|
|
140
|
+
edges[1:-1] = boundaries
|
|
141
|
+
|
|
142
|
+
# Update each centroid to the conditional expectation over its cell
|
|
143
|
+
new_centroids = np.empty_like(centroids)
|
|
144
|
+
for i in range(n_levels):
|
|
145
|
+
new_centroids[i] = _conditional_expectation(
|
|
146
|
+
edges[i], edges[i + 1], dim
|
|
147
|
+
)
|
|
148
|
+
|
|
149
|
+
delta = np.max(np.abs(new_centroids - centroids))
|
|
150
|
+
centroids = new_centroids
|
|
151
|
+
|
|
152
|
+
if delta < tol:
|
|
153
|
+
break
|
|
154
|
+
|
|
155
|
+
# Final boundaries
|
|
156
|
+
boundaries = 0.5 * (centroids[:-1] + centroids[1:])
|
|
157
|
+
return centroids.astype(np.float64), boundaries.astype(np.float64)
|
|
158
|
+
|
|
159
|
+
|
|
160
|
+
@functools.lru_cache(maxsize=256)
|
|
161
|
+
def get_codebook(
|
|
162
|
+
dim: int, bits: int
|
|
163
|
+
) -> Tuple[np.ndarray, np.ndarray]:
|
|
164
|
+
"""Return a cached Lloyd-Max codebook for the given dimension and bit-width.
|
|
165
|
+
|
|
166
|
+
On the first call for a given ``(dim, bits)`` pair the codebook is
|
|
167
|
+
computed via :func:`lloyd_max_codebook` and then cached for all
|
|
168
|
+
subsequent calls.
|
|
169
|
+
|
|
170
|
+
Parameters
|
|
171
|
+
----------
|
|
172
|
+
dim : int
|
|
173
|
+
Vector dimension.
|
|
174
|
+
bits : int
|
|
175
|
+
Bit-width per coordinate.
|
|
176
|
+
|
|
177
|
+
Returns
|
|
178
|
+
-------
|
|
179
|
+
centroids : ndarray, shape (2**bits,)
|
|
180
|
+
boundaries : ndarray, shape (2**bits - 1,)
|
|
181
|
+
"""
|
|
182
|
+
return lloyd_max_codebook(dim, bits)
|
|
@@ -0,0 +1,258 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Bit-packing and unpacking utilities.
|
|
3
|
+
|
|
4
|
+
Quantised indices (b bits each) and QJL sign bits are packed into
|
|
5
|
+
contiguous ``uint8`` arrays for compact storage.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
import numpy as np
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def pack_indices(indices: np.ndarray, bits: int) -> np.ndarray:
|
|
14
|
+
"""Pack an array of integer indices into a compact ``uint8`` byte array.
|
|
15
|
+
|
|
16
|
+
Each index value must be in ``[0, 2**bits)``. The indices are stored
|
|
17
|
+
in a flat bit-stream, packed from the LSB of each output byte.
|
|
18
|
+
|
|
19
|
+
Parameters
|
|
20
|
+
----------
|
|
21
|
+
indices : ndarray of int
|
|
22
|
+
Flat array of index values.
|
|
23
|
+
bits : int
|
|
24
|
+
Number of bits per index (1–8).
|
|
25
|
+
|
|
26
|
+
Returns
|
|
27
|
+
-------
|
|
28
|
+
packed : ndarray of uint8
|
|
29
|
+
Byte array containing the packed bits.
|
|
30
|
+
"""
|
|
31
|
+
indices = np.asarray(indices, dtype=np.uint32).ravel()
|
|
32
|
+
n = indices.size
|
|
33
|
+
total_bits = n * bits
|
|
34
|
+
total_bytes = (total_bits + 7) // 8
|
|
35
|
+
|
|
36
|
+
packed = np.zeros(total_bytes, dtype=np.uint8)
|
|
37
|
+
|
|
38
|
+
# Build a flat bit stream
|
|
39
|
+
for b in range(bits):
|
|
40
|
+
# Extract bit `b` from each index
|
|
41
|
+
bit_values = ((indices >> b) & 1).astype(np.uint8)
|
|
42
|
+
for i, bv in enumerate(bit_values):
|
|
43
|
+
flat_bit_pos = i * bits + b
|
|
44
|
+
byte_idx = flat_bit_pos // 8
|
|
45
|
+
bit_idx = flat_bit_pos % 8
|
|
46
|
+
packed[byte_idx] |= bv << bit_idx
|
|
47
|
+
|
|
48
|
+
return packed
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
def unpack_indices(packed: np.ndarray, bits: int, count: int) -> np.ndarray:
|
|
52
|
+
"""Unpack integer indices from a packed ``uint8`` byte array.
|
|
53
|
+
|
|
54
|
+
Parameters
|
|
55
|
+
----------
|
|
56
|
+
packed : ndarray of uint8
|
|
57
|
+
Byte array produced by :func:`pack_indices`.
|
|
58
|
+
bits : int
|
|
59
|
+
Number of bits per index (1–8).
|
|
60
|
+
count : int
|
|
61
|
+
Number of indices to unpack.
|
|
62
|
+
|
|
63
|
+
Returns
|
|
64
|
+
-------
|
|
65
|
+
indices : ndarray of uint32
|
|
66
|
+
Recovered index values.
|
|
67
|
+
"""
|
|
68
|
+
packed = np.asarray(packed, dtype=np.uint8).ravel()
|
|
69
|
+
indices = np.zeros(count, dtype=np.uint32)
|
|
70
|
+
|
|
71
|
+
for i in range(count):
|
|
72
|
+
val = np.uint32(0)
|
|
73
|
+
for b in range(bits):
|
|
74
|
+
flat_bit_pos = i * bits + b
|
|
75
|
+
byte_idx = flat_bit_pos // 8
|
|
76
|
+
bit_idx = flat_bit_pos % 8
|
|
77
|
+
bit_val = (packed[byte_idx] >> bit_idx) & 1
|
|
78
|
+
val |= np.uint32(bit_val) << np.uint32(b)
|
|
79
|
+
indices[i] = val
|
|
80
|
+
|
|
81
|
+
return indices
|
|
82
|
+
|
|
83
|
+
|
|
84
|
+
def pack_indices_fast(indices: np.ndarray, bits: int) -> np.ndarray:
|
|
85
|
+
"""Vectorised bit-packing for common bit-widths.
|
|
86
|
+
|
|
87
|
+
Falls back to :func:`pack_indices` for uncommon widths.
|
|
88
|
+
|
|
89
|
+
Parameters
|
|
90
|
+
----------
|
|
91
|
+
indices : ndarray of int
|
|
92
|
+
Flat array of index values.
|
|
93
|
+
bits : int
|
|
94
|
+
Number of bits per index.
|
|
95
|
+
|
|
96
|
+
Returns
|
|
97
|
+
-------
|
|
98
|
+
packed : ndarray of uint8
|
|
99
|
+
"""
|
|
100
|
+
indices = np.asarray(indices, dtype=np.uint8).ravel()
|
|
101
|
+
|
|
102
|
+
if bits == 8:
|
|
103
|
+
return indices.copy()
|
|
104
|
+
|
|
105
|
+
if bits == 4:
|
|
106
|
+
# Pack two 4-bit values per byte
|
|
107
|
+
n = indices.size
|
|
108
|
+
padded = indices
|
|
109
|
+
if n % 2 != 0:
|
|
110
|
+
padded = np.append(indices, np.uint8(0))
|
|
111
|
+
low = padded[0::2] & 0x0F
|
|
112
|
+
high = (padded[1::2] & 0x0F) << 4
|
|
113
|
+
return (low | high).astype(np.uint8)
|
|
114
|
+
|
|
115
|
+
if bits == 2:
|
|
116
|
+
# Pack four 2-bit values per byte
|
|
117
|
+
n = indices.size
|
|
118
|
+
pad_len = (4 - n % 4) % 4
|
|
119
|
+
if pad_len:
|
|
120
|
+
padded = np.append(indices, np.zeros(pad_len, dtype=np.uint8))
|
|
121
|
+
else:
|
|
122
|
+
padded = indices
|
|
123
|
+
a = padded[0::4] & 0x03
|
|
124
|
+
b = (padded[1::4] & 0x03) << 2
|
|
125
|
+
c = (padded[2::4] & 0x03) << 4
|
|
126
|
+
d = (padded[3::4] & 0x03) << 6
|
|
127
|
+
return (a | b | c | d).astype(np.uint8)
|
|
128
|
+
|
|
129
|
+
if bits == 1:
|
|
130
|
+
return pack_signs_raw(indices)
|
|
131
|
+
|
|
132
|
+
# General fallback
|
|
133
|
+
return pack_indices(indices.astype(np.uint32), bits)
|
|
134
|
+
|
|
135
|
+
|
|
136
|
+
def unpack_indices_fast(
|
|
137
|
+
packed: np.ndarray, bits: int, count: int
|
|
138
|
+
) -> np.ndarray:
|
|
139
|
+
"""Vectorised unpacking for common bit-widths.
|
|
140
|
+
|
|
141
|
+
Parameters
|
|
142
|
+
----------
|
|
143
|
+
packed : ndarray of uint8
|
|
144
|
+
bits : int
|
|
145
|
+
count : int
|
|
146
|
+
|
|
147
|
+
Returns
|
|
148
|
+
-------
|
|
149
|
+
indices : ndarray of uint32
|
|
150
|
+
"""
|
|
151
|
+
packed = np.asarray(packed, dtype=np.uint8).ravel()
|
|
152
|
+
|
|
153
|
+
if bits == 8:
|
|
154
|
+
return packed[:count].astype(np.uint32)
|
|
155
|
+
|
|
156
|
+
if bits == 4:
|
|
157
|
+
low = packed & 0x0F
|
|
158
|
+
high = (packed >> 4) & 0x0F
|
|
159
|
+
interleaved = np.empty(packed.size * 2, dtype=np.uint32)
|
|
160
|
+
interleaved[0::2] = low
|
|
161
|
+
interleaved[1::2] = high
|
|
162
|
+
return interleaved[:count]
|
|
163
|
+
|
|
164
|
+
if bits == 2:
|
|
165
|
+
a = packed & 0x03
|
|
166
|
+
b = (packed >> 2) & 0x03
|
|
167
|
+
c = (packed >> 4) & 0x03
|
|
168
|
+
d = (packed >> 6) & 0x03
|
|
169
|
+
interleaved = np.empty(packed.size * 4, dtype=np.uint32)
|
|
170
|
+
interleaved[0::4] = a
|
|
171
|
+
interleaved[1::4] = b
|
|
172
|
+
interleaved[2::4] = c
|
|
173
|
+
interleaved[3::4] = d
|
|
174
|
+
return interleaved[:count]
|
|
175
|
+
|
|
176
|
+
if bits == 1:
|
|
177
|
+
signs_raw = unpack_signs_raw(packed, count)
|
|
178
|
+
return signs_raw.astype(np.uint32)
|
|
179
|
+
|
|
180
|
+
return unpack_indices(packed, bits, count)
|
|
181
|
+
|
|
182
|
+
|
|
183
|
+
def pack_signs_raw(bits_array: np.ndarray) -> np.ndarray:
|
|
184
|
+
"""Pack a {0,1} bit array into bytes.
|
|
185
|
+
|
|
186
|
+
Parameters
|
|
187
|
+
----------
|
|
188
|
+
bits_array : ndarray
|
|
189
|
+
Array of 0/1 values.
|
|
190
|
+
|
|
191
|
+
Returns
|
|
192
|
+
-------
|
|
193
|
+
packed : ndarray of uint8
|
|
194
|
+
"""
|
|
195
|
+
bits_array = np.asarray(bits_array, dtype=np.uint8).ravel()
|
|
196
|
+
n = bits_array.size
|
|
197
|
+
pad_len = (8 - n % 8) % 8
|
|
198
|
+
if pad_len:
|
|
199
|
+
bits_array = np.append(
|
|
200
|
+
bits_array, np.zeros(pad_len, dtype=np.uint8)
|
|
201
|
+
)
|
|
202
|
+
reshaped = bits_array.reshape(-1, 8)
|
|
203
|
+
multipliers = (1 << np.arange(8, dtype=np.uint8))[np.newaxis, :]
|
|
204
|
+
return (reshaped * multipliers).sum(axis=1).astype(np.uint8)
|
|
205
|
+
|
|
206
|
+
|
|
207
|
+
def unpack_signs_raw(packed: np.ndarray, count: int) -> np.ndarray:
|
|
208
|
+
"""Unpack bytes into a {0,1} bit array.
|
|
209
|
+
|
|
210
|
+
Parameters
|
|
211
|
+
----------
|
|
212
|
+
packed : ndarray of uint8
|
|
213
|
+
count : int
|
|
214
|
+
|
|
215
|
+
Returns
|
|
216
|
+
-------
|
|
217
|
+
bits_array : ndarray of uint8
|
|
218
|
+
Array of 0/1 values.
|
|
219
|
+
"""
|
|
220
|
+
packed = np.asarray(packed, dtype=np.uint8).ravel()
|
|
221
|
+
bits = np.unpackbits(packed, bitorder="little")
|
|
222
|
+
return bits[:count]
|
|
223
|
+
|
|
224
|
+
|
|
225
|
+
def pack_signs(signs: np.ndarray) -> np.ndarray:
|
|
226
|
+
"""Pack a {-1, +1} sign array into bytes.
|
|
227
|
+
|
|
228
|
+
+1 is stored as bit 1, -1 as bit 0.
|
|
229
|
+
|
|
230
|
+
Parameters
|
|
231
|
+
----------
|
|
232
|
+
signs : ndarray
|
|
233
|
+
Array of -1 / +1 values.
|
|
234
|
+
|
|
235
|
+
Returns
|
|
236
|
+
-------
|
|
237
|
+
packed : ndarray of uint8
|
|
238
|
+
"""
|
|
239
|
+
bits = ((np.asarray(signs).ravel() + 1) // 2).astype(np.uint8)
|
|
240
|
+
return pack_signs_raw(bits)
|
|
241
|
+
|
|
242
|
+
|
|
243
|
+
def unpack_signs(packed: np.ndarray, count: int) -> np.ndarray:
|
|
244
|
+
"""Unpack bytes into a {-1, +1} sign array.
|
|
245
|
+
|
|
246
|
+
Parameters
|
|
247
|
+
----------
|
|
248
|
+
packed : ndarray of uint8
|
|
249
|
+
count : int
|
|
250
|
+
Number of sign values to recover.
|
|
251
|
+
|
|
252
|
+
Returns
|
|
253
|
+
-------
|
|
254
|
+
signs : ndarray of float64
|
|
255
|
+
Array of -1.0 / +1.0 values.
|
|
256
|
+
"""
|
|
257
|
+
bits = unpack_signs_raw(packed, count)
|
|
258
|
+
return bits.astype(np.float64) * 2.0 - 1.0
|
pyturboquant_cpu/qjl.py
ADDED
|
@@ -0,0 +1,86 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Quantized Johnson–Lindenstrauss (QJL) 1-bit inner-product quantiser.
|
|
3
|
+
|
|
4
|
+
QJL provides an *unbiased* estimator for ⟨y, x⟩ using only 1 bit per
|
|
5
|
+
coordinate. It is used as the second stage of TurboQuant_Prod to
|
|
6
|
+
correct the inner-product bias introduced by the MSE-optimal quantiser.
|
|
7
|
+
|
|
8
|
+
Reference
|
|
9
|
+
---------
|
|
10
|
+
Zandieh et al., Definition 1 and Lemma 4 of arXiv:2504.19874.
|
|
11
|
+
"""
|
|
12
|
+
|
|
13
|
+
from __future__ import annotations
|
|
14
|
+
|
|
15
|
+
import math
|
|
16
|
+
|
|
17
|
+
import numpy as np
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
def generate_qjl_matrix(
|
|
21
|
+
dim: int, seed: int | None = None
|
|
22
|
+
) -> np.ndarray:
|
|
23
|
+
"""Generate a d×d random Gaussian matrix for QJL.
|
|
24
|
+
|
|
25
|
+
Each entry is drawn i.i.d. from N(0, 1).
|
|
26
|
+
|
|
27
|
+
Parameters
|
|
28
|
+
----------
|
|
29
|
+
dim : int
|
|
30
|
+
Matrix dimension *d*.
|
|
31
|
+
seed : int or None, optional
|
|
32
|
+
Random seed for reproducibility.
|
|
33
|
+
|
|
34
|
+
Returns
|
|
35
|
+
-------
|
|
36
|
+
S : ndarray, shape (d, d)
|
|
37
|
+
"""
|
|
38
|
+
rng = np.random.default_rng(seed)
|
|
39
|
+
return rng.standard_normal((dim, dim))
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
def qjl_quantize(
|
|
43
|
+
vectors: np.ndarray, S: np.ndarray
|
|
44
|
+
) -> np.ndarray:
|
|
45
|
+
"""Apply QJL quantisation: sign(S · x).
|
|
46
|
+
|
|
47
|
+
Parameters
|
|
48
|
+
----------
|
|
49
|
+
vectors : ndarray, shape (..., d)
|
|
50
|
+
Input vectors (should already be normalised / residuals).
|
|
51
|
+
S : ndarray, shape (d, d)
|
|
52
|
+
Random Gaussian matrix.
|
|
53
|
+
|
|
54
|
+
Returns
|
|
55
|
+
-------
|
|
56
|
+
signs : ndarray, shape (..., d)
|
|
57
|
+
Array of -1.0 / +1.0 values.
|
|
58
|
+
"""
|
|
59
|
+
# vectors @ S.T gives (..., d)
|
|
60
|
+
projected = vectors @ S.T
|
|
61
|
+
signs = np.sign(projected)
|
|
62
|
+
# Replace exact zeros (extremely unlikely) with +1
|
|
63
|
+
signs[signs == 0] = 1.0
|
|
64
|
+
return signs
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
def qjl_dequantize(
|
|
68
|
+
signs: np.ndarray, S: np.ndarray
|
|
69
|
+
) -> np.ndarray:
|
|
70
|
+
"""QJL inverse map: Q_qjl^{-1}(z) = √(π/2) · S^T · z / d.
|
|
71
|
+
|
|
72
|
+
Parameters
|
|
73
|
+
----------
|
|
74
|
+
signs : ndarray, shape (..., d)
|
|
75
|
+
Array of -1 / +1 values.
|
|
76
|
+
S : ndarray, shape (d, d)
|
|
77
|
+
The same random matrix used during quantisation.
|
|
78
|
+
|
|
79
|
+
Returns
|
|
80
|
+
-------
|
|
81
|
+
reconstructed : ndarray, shape (..., d)
|
|
82
|
+
"""
|
|
83
|
+
d = S.shape[0]
|
|
84
|
+
scale = math.sqrt(math.pi / 2.0) / d
|
|
85
|
+
# signs @ S gives (..., d)
|
|
86
|
+
return scale * (signs @ S)
|