SabaPivot's picture
Add qSNU oracle scaling audit and compact index
709ff03 verified
Raw
History Blame Contribute Delete
5.2 kB
"""Theorem 5.1 exact oracles (spec: specs/exp00_core.md, Algorithms 2-4, verbatim)."""
import numpy as np
from .swf import NEG_INF, norm_q
def _softmax(x):
m = np.max(x)
e = np.exp(x - m)
return e / np.sum(e)
def oracle_wpm(u, w, q, k):
n = len(u)
q = norm_q(q)
if q == NEG_INF:
r = 1.0 / u
elif q == 0.0:
r = w.copy()
elif q == 1.0:
p = np.zeros(n)
top = np.argsort(-(w * u), kind="stable")[:k]
p[top] = 1.0
return p
else:
log_rates = (np.log(w) + q * np.log(u)) / (1.0 - q)
r = _softmax(log_rates)
order = np.argsort(-r, kind="stable") # descending
r_sorted = r[order]
R = np.cumsum(r_sorted[::-1])[::-1] # suffix sums, R[i] = sum(r_sorted[i:])
p = np.zeros(n)
rem = float(k)
for i in range(n):
t = 1.0 / r_sorted[i]
if t * R[i] > rem + 1e-15:
t_final = rem / R[i]
p[order[i:]] = t_final * r_sorted[i:]
rem = 0.0
break
p[order[i]] = 1.0
rem -= 1.0
return p
def oracle_kolm(u, w, q, k):
n = len(u)
q = norm_q(q)
if q == 0.0:
p = np.zeros(n)
top = np.argsort(-(w * u), kind="stable")[:k]
p[top] = 1.0
return p
if q == NEG_INF:
r = 1.0 / u
t_start = np.zeros(n)
t_end = u.copy()
else:
r = 1.0 / (-q * u)
t_start = -np.log(w * u)
t_end = -q * u - np.log(w * u)
events = sorted(
[(t_start[i], 1, i) for i in range(n)] + [(t_end[i], -1, i) for i in range(n)],
key=lambda e: e[0],
)
tol = 1e-9
groups = []
i = 0
ne = len(events)
while i < ne:
tau = events[i][0]
members = []
while i < ne and events[i][0] <= tau + tol:
members.append((events[i][1], events[i][2]))
i += 1
groups.append((tau, members))
p = np.zeros(n)
active = set()
t_prev = groups[0][0]
for tau, members in groups:
dt = tau - t_prev
if dt > 0 and active:
rate_sum = sum(r[j] for j in active)
m = rate_sum * dt
if p.sum() + m > k + 1e-12:
dt_f = (k - p.sum()) / rate_sum
for j in active:
p[j] += r[j] * dt_f
return p
for j in active:
p[j] += r[j] * dt
for kind, j in members:
if kind == 1:
active.add(j)
else:
active.discard(j)
t_prev = tau
return p
def oracle_gini(u, w, k):
# Blocks B in Algorithm 4 are not simply the single remaining suffix: an
# entry's fill priority is its standalone rate e_i = w_i*u_i, and a block
# must be filled to a common (tied) level whenever a lower-u member would
# otherwise be outpaced by a higher-u, higher-priority neighbor (breaking
# the required non-decreasing v-in-u-order). Build blocks via PAVA: scan
# u-ascending, merging a new (larger-u) singleton into the block
# immediately to its left whenever that would violate a non-decreasing
# left-to-right r(B) -- the resulting blocks have r(B) non-decreasing
# left-to-right, i.e. strictly decreasing in *fill priority* order, and
# are filled highest-rate-block first.
n = len(u)
order = np.argsort(u) # ascending: u_(i) = i-th smallest
u_sorted = u[order]
w_sorted = w
tol = 1e-12
stack = [] # list of [l, r, sum_w, sum_invu], appended smallest-u-first
for i in range(n):
l, r = i, i
sum_w = w_sorted[i]
sum_invu = 1.0 / u_sorted[i]
while stack and sum_w * stack[-1][3] <= stack[-1][2] * sum_invu + tol:
pl, pr, pw, pinvu = stack.pop()
l = pl
sum_w += pw
sum_invu += pinvu
stack.append([l, r, sum_w, sum_invu])
stack.sort(key=lambda b: b[2] / b[3], reverse=True) # fill by decreasing r(B)
p_sorted = np.zeros(n)
rem = float(k)
for l0, r0, _, _ in stack:
if rem <= tol:
break
ll = l0
while rem > tol and ll <= r0:
invsum = float(np.sum(1.0 / u_sorted[ll:r0 + 1]))
dt = (1.0 - p_sorted[ll]) * u_sorted[ll]
m = invsum * dt
if m <= rem + 1e-12:
p_sorted[ll:r0 + 1] += dt / u_sorted[ll:r0 + 1]
rem -= m
ll += 1
else:
p_sorted[ll:r0 + 1] += rem / (invsum * u_sorted[ll:r0 + 1])
rem = 0.0
p = np.zeros(n)
p[order] = p_sorted
return p
def oracle(u, family, w, q, k):
"""argmax_{p in P_k} M(u . p), P_k = {p in [0,1]^n : sum p_i = k}."""
u = np.asarray(u, dtype=float)
w = np.asarray(w, dtype=float)
assert np.all(u > 0), "oracle requires u_i > 0"
if family == "wpm":
p = oracle_wpm(u, w, q, k)
elif family == "kolm":
p = oracle_kolm(u, w, q, k)
elif family == "gini":
p = oracle_gini(u, w, k)
else:
raise ValueError(f"unknown family {family!r}")
p = np.clip(p, 0.0, 1.0)
assert abs(p.sum() - k) < 1e-9, f"oracle sum {p.sum()} != k {k}"
return p