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,
)
| 18 | |
| 19 | |
| 20 | def 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 | |
| 61 | class LloydMaxCodebook: |
no test coverage detected