summaryrefslogtreecommitdiff
path: root/src/zbp_scaling/zbp/probes.py
diff options
context:
space:
mode:
authoryurenh <blackhao0426@gmail.com>2026-08-31 18:14:09 -0500
committeryurenh <blackhao0426@gmail.com>2026-08-31 18:14:09 -0500
commit6a544fabfc2af22e4d5823410dd2387b5af89ea9 (patch)
tree0abd67bdda420deed27428b621fb59db8be07f41 /src/zbp_scaling/zbp/probes.py
scaffold: model (OLMo2-ish + ZBP partition), trainer (DDP/config), data shards, bench
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01GkgLsACEF6CCP7EUfA5fZe
Diffstat (limited to 'src/zbp_scaling/zbp/probes.py')
-rw-r--r--src/zbp_scaling/zbp/probes.py78
1 files changed, 78 insertions, 0 deletions
diff --git a/src/zbp_scaling/zbp/probes.py b/src/zbp_scaling/zbp/probes.py
new file mode 100644
index 0000000..a8e7fe6
--- /dev/null
+++ b/src/zbp_scaling/zbp/probes.py
@@ -0,0 +1,78 @@
+"""Random probe directions with E[u u^T] = I_d per sample unit.
+
+All families return entries of O(1) magnitude so that `x + eps * u` moves every
+coordinate by ~eps. Divide by sqrt(d) for the fixed-Euclidean-norm variant.
+"""
+import math
+import torch
+
+_HADAMARD_CACHE = {}
+
+
+def _hadamard(n_pow2, device, dtype):
+ key = (n_pow2, str(device), dtype)
+ H = _HADAMARD_CACHE.get(key)
+ if H is None:
+ H = torch.ones(1, 1, dtype=dtype, device=device)
+ while H.shape[0] < n_pow2:
+ H = torch.cat([torch.cat([H, H], 1), torch.cat([H, -H], 1)], 0)
+ _HADAMARD_CACHE[key] = H
+ return H
+
+
+def sample_probes(n, batch_shape, d, family, gen, device, dtype, cycle=None):
+ """Return u of shape [n, *batch_shape, d].
+
+ For every sample unit (index into batch_shape) the n probes satisfy
+ E[u_i u_i^T] = I_d. 'orthogonal' and 'hadamard' additionally make the n probes of
+ one sample unit mutually orthogonal (n <= d).
+ """
+ B = math.prod(batch_shape) if len(batch_shape) else 1
+ if family == "rademacher":
+ u = torch.randint(0, 2, (n, B, d), generator=gen, device=device).to(dtype) * 2 - 1
+ elif family == "gaussian":
+ u = torch.randn((n, B, d), generator=gen, device=device, dtype=dtype)
+ elif family == "orthogonal":
+ if n > d:
+ raise ValueError(f"orthogonal probes need n <= d (n={n}, d={d})")
+ A = torch.randn((B, d, n), generator=gen, device=device, dtype=dtype)
+ Q, R = torch.linalg.qr(A) # [B, d, n], orthonormal columns
+ s = torch.sign(torch.diagonal(R, dim1=-2, dim2=-1)) # make the distribution Haar
+ s[s == 0] = 1
+ Q = Q * s[:, None, :]
+ u = Q.permute(2, 0, 1) * math.sqrt(d) # columns are unit vectors -> scale by sqrt(d)
+ elif family == "hadamard":
+ d2 = 1 << (d - 1).bit_length()
+ if n > d2:
+ raise ValueError(f"hadamard probes need n <= {d2}")
+ H = _hadamard(d2, device, dtype) # +-1 entries, orthogonal rows
+ signs = torch.randint(0, 2, (B, d2), generator=gen, device=device).to(dtype) * 2 - 1
+ rows = torch.argsort(torch.rand((B, d2), generator=gen, device=device), dim=1)[:, :n]
+ u = H[rows] * signs[:, None, :] # [B, n, d2]
+ u = u[:, :, :d].permute(1, 0, 2).contiguous()
+ elif family == "hadamard_cycle":
+ # Temporal QMC: every unit of a step shares the same n Hadamard rows (advancing by n per step) and the same
+ # sign vector (fixed for a whole cycle of d2/n steps), so the projection noise cancels exactly over a cycle
+ # for a slowly varying gradient -- at the price of no cross-unit averaging within a step.
+ d2 = 1 << (d - 1).bit_length()
+ H = _hadamard(d2, device, dtype)
+ step = int(cycle or 0)
+ start, cyc = (step * n) % d2, (step * n) // d2
+ rows = (start + torch.arange(n, device=device)) % d2
+ g2 = torch.Generator(device=device)
+ g2.manual_seed(1000003 * cyc + 17)
+ signs = torch.randint(0, 2, (1, d2), generator=g2, device=device).to(dtype) * 2 - 1
+ u = (H[rows] * signs)[:, None, :d].expand(n, B, d).contiguous()
+ elif family == "coordinate_all":
+ # deterministic sweep of every coordinate (requires n == d): u_i = sqrt(d) e_i, so (1/n) sum u_i u_i^T g = g
+ # exactly -- the coordinate-finite-difference estimator of HZO / BOND at its full budget of 2d queries per unit
+ if n != d:
+ raise ValueError(f"coordinate_all needs n == d (n={n}, d={d})")
+ u = (math.sqrt(d) * torch.eye(d, device=device, dtype=dtype))[:, None, :].expand(n, B, d).contiguous()
+ elif family == "coordinate":
+ idx = torch.randint(0, d, (n, B), generator=gen, device=device)
+ u = torch.zeros((n, B, d), device=device, dtype=dtype)
+ u.scatter_(2, idx[..., None], math.sqrt(d))
+ else:
+ raise ValueError(f"unknown probe family {family!r}")
+ return u.reshape(n, *batch_shape, d)