MCPcopy Create free account
hub / github.com/DevTechJr/turboquant_cutile / solve_lloyd_max

Function solve_lloyd_max

turboquant_cutile/codebook.py:20–58  ·  view source on GitHub ↗

Returns (centroids, boundaries) as sorted float32 tensors. centroids: (2^bits,) boundaries: (2^bits - 1,)

(
    d: int,
    bits: int,
    max_iter: int = 200,
    tol: float = 1e-10,
)

Source from the content-addressed store, hash-verified

18
19
20def solve_lloyd_max(
21 d: int,
22 bits: int,
23 max_iter: int = 200,
24 tol: float = 1e-10,
25) -> tuple[torch.Tensor, torch.Tensor]:
26 """
27 Returns (centroids, boundaries) as sorted float32 tensors.
28 centroids: (2^bits,) boundaries: (2^bits - 1,)
29 """
30 n_levels = 1 << bits
31 sigma = 1.0 / math.sqrt(d)
32 pdf = lambda x: _gaussian_pdf(x, sigma)
33
34 lo, hi = -3.5 * sigma, 3.5 * sigma
35 centroids = [lo + (hi - lo) * (i + 0.5) / n_levels for i in range(n_levels)]
36
37 for _ in range(max_iter):
38 boundaries = [
39 (centroids[i] + centroids[i + 1]) / 2.0 for i in range(n_levels - 1)
40 ]
41 edges = [lo * 3] + boundaries + [hi * 3]
42 new_centroids = []
43 for i in range(n_levels):
44 a, b = edges[i], edges[i + 1]
45 num, _ = integrate.quad(lambda x: x * pdf(x), a, b)
46 den, _ = integrate.quad(pdf, a, b)
47 new_centroids.append(num / den if den > 1e-15 else centroids[i])
48 if max(abs(new_centroids[i] - centroids[i]) for i in range(n_levels)) < tol:
49 break
50 centroids = new_centroids
51
52 boundaries = [
53 (centroids[i] + centroids[i + 1]) / 2.0 for i in range(n_levels - 1)
54 ]
55 return (
56 torch.tensor(centroids, dtype=torch.float32),
57 torch.tensor(boundaries, dtype=torch.float32),
58 )
59
60
61class LloydMaxCodebook:

Callers 1

__init__Method · 0.85

Calls 1

_gaussian_pdfFunction · 0.85

Tested by

no test coverage detected