import numpy as np from remedi.optimization import solve_regret, regret_components, softmin, finite_minimax def test_factorized_comparator_scan_matches_dense_covariance(): rng = np.random.default_rng(1) b = rng.normal(size=(7,3)); d = rng.uniform(size=7) p = rng.dirichlet(np.ones(7)); mean = rng.normal(size=7) c = b@b.T+np.diag(d) expected = [] for e in np.eye(7): z = p-e expected.append(z@mean+1.2*np.sqrt(z@c@z)) np.testing.assert_allclose(regret_components(p,mean,b,d,1.2), expected, atol=1e-10) def test_zero_radius_has_analytic_gibbs_solution(): mean = np.array([1.,2.,.5]); prior = np.array([.2,.3,.5]) result = solve_regret(mean,np.zeros((3,1)),np.ones(3),beta=0,tau=.2,prior=prior) np.testing.assert_allclose(result.probabilities,softmin(mean,.2,prior)) def test_common_loss_offsets_and_common_error_factors_cancel(): mean = np.array([.2,.3,.8]); b = np.array([[.1],[.2],[.4]]); d = np.full(3,.01) a = solve_regret(mean,b,d,tau=.1) c = solve_regret(mean+123,np.column_stack([b,np.full(3,7.)]),d,tau=.1) np.testing.assert_allclose(a.probabilities,c.probabilities,atol=2e-4) def test_constraint_generation_matches_full_program(): rng = np.random.default_rng(8) mean=rng.normal(size=10); b=rng.normal(size=(10,3))*.2; d=np.full(10,.02) full=solve_regret(mean,b,d,tau=.2,constraint_generation=False) active=solve_regret(mean,b,d,tau=.2) np.testing.assert_allclose(full.probabilities,active.probabilities,atol=3e-4) assert active.constraint_violation <= 1e-6 def test_conditional_bound_contains_sampled_loss_vectors(): rng=np.random.default_rng(3); mean=rng.normal(size=5); b=rng.normal(size=(5,2)); d=np.full(5,.05) result=solve_regret(mean,b,d,beta=1.5,tau=.1) factor=np.column_stack([b,np.diag(np.sqrt(d))]) for _ in range(30): u=rng.normal(size=7);u=u/np.linalg.norm(u)*rng.uniform(0,1.5) loss=mean+factor@u assert result.probabilities@loss-loss.min() <= result.worst_regret+1e-6 def test_finite_scenario_minimax_is_offset_invariant(): losses=np.array([[0.,3.,2.],[3.,0.,2.]]) a=finite_minimax(losses,tau=.1) b=finite_minimax(losses+np.array([[100.],[-20.]]),tau=.1) np.testing.assert_allclose(a,b,atol=1e-5)