| """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") |
| r_sorted = r[order] |
| R = np.cumsum(r_sorted[::-1])[::-1] |
|
|
| 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): |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| n = len(u) |
| order = np.argsort(u) |
| u_sorted = u[order] |
| w_sorted = w |
| tol = 1e-12 |
|
|
| stack = [] |
| 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) |
|
|
| 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 |
|
|