ReMEDi / src /remedi /optimization.py
pranamanam's picture
Upload 62 files
3f98d52 verified
Raw
History Blame Contribute Delete
7.08 kB
"""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}")