File size: 7,081 Bytes
3f98d52
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
"""Convex finite-action minimax regret and matched generation rules."""
from dataclasses import dataclass
import time
import numpy as np
from scipy.special import logsumexp

@dataclass
class Decision:
    probabilities: np.ndarray
    objective: float
    worst_regret: float
    constraint_violation: float
    seconds: float
    status: str

def validate_problem(mean, factor, diagonal, beta, tau, prior=None):
    mean = np.asarray(mean, dtype=float)
    factor = np.asarray(factor, dtype=float)
    diagonal = np.asarray(diagonal, dtype=float)
    n = len(mean)
    if n == 0 or mean.shape != (n,) or factor.ndim != 2 or factor.shape[0] != n or diagonal.shape != (n,):
        raise ValueError("Incompatible finite-action dimensions")
    if not all(np.isfinite(x).all() for x in [mean, factor, diagonal]):
        raise ValueError("Nonfinite loss or uncertainty values")
    if (diagonal < 0).any() or not np.isfinite(beta) or not np.isfinite(tau) or beta < 0 or tau <= 0:
        raise ValueError("Need nonnegative uncertainty and positive temperature")
    prior = np.ones(n)/n if prior is None else np.asarray(prior, dtype=float)
    if prior.shape != (n,) or not np.isfinite(prior).all() or (prior <= 0).any():
        raise ValueError("Every eligible action needs a positive finite prior")
    return mean, factor, diagonal, prior/prior.sum()

def regret_components(p, mean, factor, diagonal, beta):
    """All comparator constraints in O(N*S), including diagonal cross-terms."""
    q = factor.T @ p
    variance = np.sum((factor-q[None, :])**2, axis=1)
    variance += np.dot(diagonal, p*p)-2*diagonal*p+diagonal
    return np.dot(p, mean)-mean+beta*np.sqrt(np.maximum(variance, 0))

def softmin(values, tau, prior=None):
    values = np.asarray(values, float)
    if tau <= 0 or not np.isfinite(values).all():
        raise ValueError("Softmin needs finite values and positive temperature")
    prior = np.ones(len(values))/len(values) if prior is None else np.asarray(prior, float)
    if (prior <= 0).any():
        raise ValueError("Prior must have positive support")
    logits = np.log(prior)-values/tau
    return np.exp(logits-logsumexp(logits))

def solve_regret(mean, factor, diagonal, beta=1., tau=.05, prior=None,
                 tolerance=1e-6, constraint_generation=True, max_rounds=100, solver="CLARABEL"):
    """Optimize a distribution over N candidate molecule-dose pairs.

    mean is (N,), factor is (N, S), and diagonal is (N,). Their covariance
    is factor @ factor.T + diag(diagonal). beta sets the uncertainty radius
    and tau sets KL regularization. Return probabilities and solver diagnostics.
    """
    import cvxpy as cp
    start = time.perf_counter()
    mean, factor, diagonal, prior = validate_problem(mean, factor, diagonal, beta, tau, prior)
    n = len(mean)
    # At zero uncertainty, the optimizer is the prior-weighted Gibbs distribution.
    if beta == 0:
        p = softmin(mean, tau, prior)
        regret = float(p@mean-mean.min())
        objective = regret+tau*float(np.sum(p*np.log(np.maximum(p, 1e-300)/prior)))
        return Decision(p, objective, regret, 0., time.perf_counter()-start, "analytic")
    active = set(range(n)) if not constraint_generation else {int(np.argmin(mean))}
    probability, epigraph = cp.Variable(n), cp.Variable()
    p = None
    for _ in range(max_rounds):
        constraints = [probability >= 0, cp.sum(probability) == 1]
        for b in sorted(active):
            e = np.zeros(n); e[b] = 1
            # Compare the randomized intervention with the single action b.
            contrast = probability-e
            norm = cp.norm(cp.hstack([factor.T@contrast, cp.multiply(np.sqrt(diagonal), contrast)]), 2)
            constraints.append(epigraph >= contrast@mean+beta*norm)
        problem = cp.Problem(cp.Minimize(epigraph+tau*cp.sum(cp.kl_div(probability, prior))), constraints)
        options = {"tol_gap_abs": tolerance*.1, "tol_feas": tolerance*.1} if solver == "CLARABEL" else {}
        problem.solve(solver=solver, warm_start=True, **options)
        if problem.status not in {"optimal", "optimal_inaccurate"} or probability.value is None:
            raise RuntimeError(f"Cone solver failed: {problem.status}")
        p = np.maximum(np.asarray(probability.value).ravel(), 0)
        p /= p.sum()
        # Check every comparator, including those outside the active constraint set.
        values = regret_components(p, mean, factor, diagonal, beta)
        worst = int(np.argmax(values))
        violation = max(0., float(values[worst]-epigraph.value))
        if violation <= tolerance:
            status = problem.status
            break
        # Add the most violated comparator and solve the expanded program.
        active.add(worst)
    else:
        status = "iteration_limit"
    robust_regret = float(np.max(regret_components(p, mean, factor, diagonal, beta)))
    kl = float(np.sum(p*np.log(np.maximum(p, 1e-300)/prior)))
    return Decision(p, robust_regret+tau*kl, robust_regret, violation,
                    time.perf_counter()-start, status)

def finite_minimax(losses, tau=.05, prior=None, solver="CLARABEL"):
    import cvxpy as cp
    losses = np.asarray(losses, float)
    if losses.ndim != 2 or not np.isfinite(losses).all():
        raise ValueError("Scenarios must be a finite S-by-N matrix")
    n = losses.shape[1]
    prior = np.ones(n)/n if prior is None else np.asarray(prior, float)
    regrets = losses-losses.min(axis=1, keepdims=True)
    p, z = cp.Variable(n), cp.Variable()
    problem = cp.Problem(cp.Minimize(z+tau*cp.sum(cp.kl_div(p, prior))),
                         [p >= 0, cp.sum(p) == 1, regrets@p <= z])
    problem.solve(solver=solver)
    if p.value is None:
        raise RuntimeError(f"Finite minimax failed: {problem.status}")
    values = np.maximum(np.asarray(p.value), 0)
    return values/values.sum()

def generate_baseline(name, mean, factor, diagonal, losses, beta=1., tau=.05, seed=0):
    n = len(mean)
    if name == "uniform":
        return np.ones(n)/n
    if name == "mean":
        p = np.zeros(n); p[int(np.argmin(mean))] = 1; return p
    if name == "gibbs":
        return softmin(mean, tau)
    if name == "marginal":
        return softmin(mean+beta*np.sqrt((factor**2).sum(1)+diagonal), tau)
    if name == "thompson":
        # Exact empirical distribution of scenario-wise minimizers.
        return np.bincount(np.argmin(losses, axis=1), minlength=n)/len(losses)
    if name == "finite-minimax":
        return finite_minimax(losses, tau)
    if name == "absolute":
        import cvxpy as cp
        p = cp.Variable(n)
        norm = cp.norm(cp.hstack([factor.T@p, cp.multiply(np.sqrt(diagonal), p)]), 2)
        objective = cp.Minimize(mean@p+beta*norm+tau*cp.sum(cp.kl_div(p, np.ones(n)/n)))
        problem = cp.Problem(objective, [p >= 0, cp.sum(p) == 1]); problem.solve(solver="CLARABEL")
        if p.value is None: raise RuntimeError("Absolute-loss solver failed")
        values = np.maximum(p.value, 0); return values/values.sum()
    raise ValueError(f"Unknown baseline {name}")