File size: 4,290 Bytes
f2c0505
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
from __future__ import annotations

import math
from typing import Iterable, Tuple

import numpy as np
import torch
from scipy.special import betainc, betaincinv, gammaln

EPS = 1e-10


def largest_power_of_two_divisor(d: int) -> int:
    if d <= 0:
        raise ValueError(f"d must be positive, got {d}")
    return d & -d


def coordinate_cdf(t: np.ndarray | float, d: int) -> np.ndarray:
    x = np.asarray(t, dtype=np.float64)
    x = np.clip(x, -1.0, 1.0)
    a = (d - 1.0) / 2.0
    z = betainc(0.5, a, x * x)
    return np.where(x >= 0.0, 0.5 + 0.5 * z, 0.5 - 0.5 * z)


def coordinate_ppf(p: np.ndarray | float, d: int) -> np.ndarray:
    q = np.asarray(p, dtype=np.float64)
    q = np.clip(q, np.finfo(np.float64).eps, 1.0 - np.finfo(np.float64).eps)
    a = (d - 1.0) / 2.0
    upper = q >= 0.5
    z = np.empty_like(q)
    z[upper] = betaincinv(0.5, a, 2.0 * q[upper] - 1.0)
    z[~upper] = betaincinv(0.5, a, 1.0 - 2.0 * q[~upper])
    ans = np.sqrt(np.clip(z, 0.0, 1.0))
    ans[~upper] *= -1.0
    return ans


def interval_probability(lo: float, hi: float, d: int) -> float:
    return float(coordinate_cdf(hi, d) - coordinate_cdf(lo, d))


def interval_first_moment(lo: float, hi: float, d: int) -> float:
    a = (d - 1.0) / 2.0
    log_c = gammaln(d / 2.0) - 0.5 * math.log(math.pi) - gammaln((d - 1.0) / 2.0)
    c = math.exp(log_c)
    lterm = max(0.0, 1.0 - lo * lo) ** a
    hterm = max(0.0, 1.0 - hi * hi) ** a
    return c / (d - 1.0) * (lterm - hterm)


def lloyd_max_codebook(d: int, bits: int = 4, tol: float = 1e-13, max_iter: int = 500) -> np.ndarray:
    """Deterministic Lloyd-Max solver for OrbitQuant's exact coordinate density f_d."""
    if d < 2:
        raise ValueError("d must be >= 2")
    if bits < 1 or bits > 8:
        raise ValueError("bits must be in [1,8]")
    levels = 1 << bits
    probs = (np.arange(levels, dtype=np.float64) + 0.5) / levels
    centroids = coordinate_ppf(probs, d)
    centroids = 0.5 * (centroids - centroids[::-1])
    for _ in range(max_iter):
        edges = np.empty(levels + 1, dtype=np.float64)
        edges[0], edges[-1] = -1.0, 1.0
        edges[1:-1] = 0.5 * (centroids[:-1] + centroids[1:])
        updated = np.empty_like(centroids)
        for i in range(levels):
            lo, hi = float(edges[i]), float(edges[i + 1])
            mass = interval_probability(lo, hi, d)
            updated[i] = centroids[i] if mass <= 1e-300 else interval_first_moment(lo, hi, d) / mass
        updated = 0.5 * (updated - updated[::-1])
        if np.max(np.abs(updated - centroids)) < tol:
            centroids = updated
            break
        centroids = updated
    if not np.all(np.diff(centroids) > 0):
        raise RuntimeError(f"non-monotonic Lloyd-Max codebook for d={d}")
    return centroids.astype(np.float64)


def make_rotation(d: int, seed: int) -> Tuple[np.ndarray, np.ndarray, int]:
    # Exact deterministic construction used by this Project-A runtime.
    # The OrbitQuant paper does not publish the authors' random seed.
    ss = np.random.SeedSequence([int(seed), int(d), 0x4F524249])
    rng = np.random.default_rng(ss)
    perm = rng.permutation(d).astype(np.int64)
    signs = rng.choice(np.array([-1, 1], dtype=np.int8), size=d, replace=True)
    return perm, signs, largest_power_of_two_divisor(d)


def fwht_last_dim(x: torch.Tensor, block_size: int) -> torch.Tensor:
    if block_size == 1:
        return x
    if block_size <= 0 or block_size & (block_size - 1):
        raise ValueError(f"block_size must be power-of-two, got {block_size}")
    d = int(x.shape[-1])
    if d % block_size:
        raise ValueError(f"last dimension {d} not divisible by {block_size}")
    lead = x.shape[:-1]
    y = x.reshape(*lead, d // block_size, block_size)
    step = 1
    while step < block_size:
        shape = y.shape
        z = y.reshape(*shape[:-1], -1, 2, step)
        a, b = z[..., 0, :], z[..., 1, :]
        y = torch.stack((a + b, a - b), dim=-2).reshape(*shape)
        step *= 2
    return (y / math.sqrt(block_size)).reshape(*lead, d)


def nearest_codes(x: torch.Tensor, codebook: torch.Tensor) -> torch.Tensor:
    cb = codebook.to(device=x.device, dtype=x.dtype)
    thresholds = 0.5 * (cb[:-1] + cb[1:])
    return torch.bucketize(x.contiguous(), thresholds).to(torch.uint8)