"""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