remex 0.5.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.
- remex/__init__.py +40 -0
- remex/codebook.py +113 -0
- remex/core.py +898 -0
- remex/gpu.py +503 -0
- remex/packing.py +200 -0
- remex/rotation.py +27 -0
- remex-0.5.0.dist-info/METADATA +314 -0
- remex-0.5.0.dist-info/RECORD +11 -0
- remex-0.5.0.dist-info/WHEEL +5 -0
- remex-0.5.0.dist-info/licenses/LICENSE +21 -0
- remex-0.5.0.dist-info/top_level.txt +1 -0
remex/__init__.py
ADDED
|
@@ -0,0 +1,40 @@
|
|
|
1
|
+
"""remex: Retrieval-validated embedding compression.
|
|
2
|
+
|
|
3
|
+
Compress embeddings 2-16x with proven recall. Based on the rotation + Lloyd-Max
|
|
4
|
+
scalar quantization insight from TurboQuant (Zandieh et al., ICLR 2026,
|
|
5
|
+
arXiv:2504.19874). Implements the MSE-optimal stage which empirically
|
|
6
|
+
outperforms the full TurboQuant Prod variant for nearest-neighbor retrieval.
|
|
7
|
+
|
|
8
|
+
Supports Matryoshka bit precision: encode once at full bit-width, then
|
|
9
|
+
search at any lower precision via right-shifting indices. Enables two-stage
|
|
10
|
+
coarse-to-fine retrieval from a single encoded representation.
|
|
11
|
+
|
|
12
|
+
Formerly known as polar-embed.
|
|
13
|
+
"""
|
|
14
|
+
|
|
15
|
+
import warnings
|
|
16
|
+
|
|
17
|
+
from remex.core import Quantizer, CompressedVectors, PackedVectors
|
|
18
|
+
from remex.codebook import lloyd_max_codebook, nested_codebooks
|
|
19
|
+
from remex.packing import pack, unpack, packed_nbytes
|
|
20
|
+
|
|
21
|
+
__version__ = "0.5.0"
|
|
22
|
+
__all__ = [
|
|
23
|
+
"Quantizer", "CompressedVectors", "PackedVectors",
|
|
24
|
+
"PolarQuantizer", # deprecated alias
|
|
25
|
+
"lloyd_max_codebook", "nested_codebooks",
|
|
26
|
+
"pack", "unpack", "packed_nbytes",
|
|
27
|
+
]
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
def __getattr__(name):
|
|
31
|
+
if name == "PolarQuantizer":
|
|
32
|
+
warnings.warn(
|
|
33
|
+
"PolarQuantizer has been renamed to Quantizer. "
|
|
34
|
+
"The PolarQuantizer alias is deprecated and will be "
|
|
35
|
+
"removed in a future release.",
|
|
36
|
+
DeprecationWarning,
|
|
37
|
+
stacklevel=2,
|
|
38
|
+
)
|
|
39
|
+
return Quantizer
|
|
40
|
+
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
remex/codebook.py
ADDED
|
@@ -0,0 +1,113 @@
|
|
|
1
|
+
"""Lloyd-Max optimal scalar quantizer for post-rotation coordinate distribution."""
|
|
2
|
+
|
|
3
|
+
import numpy as np
|
|
4
|
+
from scipy.stats import norm
|
|
5
|
+
from typing import Dict, Tuple
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
def lloyd_max_codebook(
|
|
9
|
+
d: int, bits: int, n_iter: int = 300
|
|
10
|
+
) -> Tuple[np.ndarray, np.ndarray]:
|
|
11
|
+
"""
|
|
12
|
+
Build optimal Lloyd-Max codebook for N(0, 1/d) distributed coordinates.
|
|
13
|
+
|
|
14
|
+
After random orthogonal rotation of unit vectors in R^d, each coordinate
|
|
15
|
+
follows a distribution that concentrates to N(0, 1/d) as d grows.
|
|
16
|
+
The Lloyd-Max quantizer minimizes MSE for this known distribution.
|
|
17
|
+
|
|
18
|
+
Args:
|
|
19
|
+
d: Vector dimension (determines coordinate variance = 1/d).
|
|
20
|
+
bits: Quantization bit-width (produces 2^bits levels).
|
|
21
|
+
n_iter: Lloyd-Max iteration count.
|
|
22
|
+
|
|
23
|
+
Returns:
|
|
24
|
+
boundaries: (2^bits - 1,) decision boundaries for np.searchsorted.
|
|
25
|
+
centroids: (2^bits,) reconstruction values.
|
|
26
|
+
"""
|
|
27
|
+
n_levels = 2**bits
|
|
28
|
+
sigma = 1.0 / np.sqrt(d)
|
|
29
|
+
rv = norm(0, sigma)
|
|
30
|
+
|
|
31
|
+
centroids = np.linspace(-3 * sigma, 3 * sigma, n_levels)
|
|
32
|
+
|
|
33
|
+
for _ in range(n_iter):
|
|
34
|
+
bounds = np.concatenate(
|
|
35
|
+
[[-np.inf], (centroids[:-1] + centroids[1:]) / 2, [np.inf]]
|
|
36
|
+
)
|
|
37
|
+
for j in range(n_levels):
|
|
38
|
+
lo, hi = bounds[j], bounds[j + 1]
|
|
39
|
+
prob = rv.cdf(hi) - rv.cdf(lo)
|
|
40
|
+
if prob > 1e-15:
|
|
41
|
+
centroids[j] = sigma**2 * (rv.pdf(lo) - rv.pdf(hi)) / prob
|
|
42
|
+
|
|
43
|
+
boundaries = (centroids[:-1] + centroids[1:]) / 2.0
|
|
44
|
+
return boundaries.astype(np.float32), centroids.astype(np.float32)
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
def nested_codebooks(
|
|
48
|
+
d: int, max_bits: int
|
|
49
|
+
) -> Dict[int, np.ndarray]:
|
|
50
|
+
"""
|
|
51
|
+
Build nested centroid tables for Matryoshka-style bit precision.
|
|
52
|
+
|
|
53
|
+
Encodes at max_bits precision. For each coarser bit level b < max_bits,
|
|
54
|
+
derives centroids by probability-weighted grouping of the max_bits
|
|
55
|
+
centroids. The top b bits of a max_bits index are a valid b-bit index
|
|
56
|
+
into the corresponding centroid table.
|
|
57
|
+
|
|
58
|
+
The Gaussian distribution is successively refinable, so the nesting
|
|
59
|
+
penalty is small (typically <1.5% recall vs independently optimized
|
|
60
|
+
codebooks at each level).
|
|
61
|
+
|
|
62
|
+
Args:
|
|
63
|
+
d: Vector dimension.
|
|
64
|
+
max_bits: Maximum quantization bit-width.
|
|
65
|
+
|
|
66
|
+
Returns:
|
|
67
|
+
Dict mapping bit-width to centroid array:
|
|
68
|
+
{max_bits: (2^max_bits,), max_bits-1: (2^(max_bits-1),), ..., 1: (2,)}
|
|
69
|
+
"""
|
|
70
|
+
_, centroids_max = lloyd_max_codebook(d, max_bits)
|
|
71
|
+
n_max = len(centroids_max)
|
|
72
|
+
sigma = 1.0 / np.sqrt(d)
|
|
73
|
+
rv = norm(0, sigma)
|
|
74
|
+
|
|
75
|
+
# Probability mass for each max_bits bin
|
|
76
|
+
bounds_max = (centroids_max[:-1] + centroids_max[1:]) / 2.0
|
|
77
|
+
full_bounds = np.concatenate([[-np.inf], bounds_max, [np.inf]])
|
|
78
|
+
probs = np.array(
|
|
79
|
+
[rv.cdf(full_bounds[i + 1]) - rv.cdf(full_bounds[i]) for i in range(n_max)]
|
|
80
|
+
)
|
|
81
|
+
|
|
82
|
+
result = {max_bits: centroids_max}
|
|
83
|
+
|
|
84
|
+
for target_bits in range(max_bits - 1, 0, -1):
|
|
85
|
+
n_target = 2**target_bits
|
|
86
|
+
group_size = n_max // n_target
|
|
87
|
+
nested_centroids = np.empty(n_target, dtype=np.float32)
|
|
88
|
+
|
|
89
|
+
for g in range(n_target):
|
|
90
|
+
start = g * group_size
|
|
91
|
+
end = start + group_size
|
|
92
|
+
group_probs = probs[start:end]
|
|
93
|
+
total_prob = group_probs.sum()
|
|
94
|
+
if total_prob > 1e-15:
|
|
95
|
+
nested_centroids[g] = np.average(
|
|
96
|
+
centroids_max[start:end], weights=group_probs
|
|
97
|
+
)
|
|
98
|
+
else:
|
|
99
|
+
nested_centroids[g] = centroids_max[start:end].mean()
|
|
100
|
+
|
|
101
|
+
result[target_bits] = nested_centroids
|
|
102
|
+
|
|
103
|
+
return result
|
|
104
|
+
|
|
105
|
+
|
|
106
|
+
def theoretical_mse(d: int, bits: int) -> float:
|
|
107
|
+
"""Theoretical MSE upper bound from TurboQuant Theorem 1."""
|
|
108
|
+
return np.sqrt(3 * np.pi) / 2 * 4 ** (-bits)
|
|
109
|
+
|
|
110
|
+
|
|
111
|
+
def theoretical_lower_bound(bits: int) -> float:
|
|
112
|
+
"""Information-theoretic lower bound on MSE (Theorem 3)."""
|
|
113
|
+
return 4 ** (-bits)
|