squaredcuber's picture
download
raw
11.9 kB
from __future__ import annotations
import numpy as np
import pytest
from loss_aware_dro_repro.bilevel import bilevel_steps, bootstrap_gaussian_moments, evaluate_outer_state, gaussian_moments, initial_radius, stabilize_metric_factor, stopping_diagnostics
from loss_aware_dro_repro.conic import clarabel_settings_contract, solve_gaussian_value
from loss_aware_dro_repro.core import ContractError, load_plan
from loss_aware_dro_repro.datasets import generate_dataset
from loss_aware_dro_repro.gelbrich import gelbrich_value_gradient, lower_triangular_central_difference, relative_gradient_error
from loss_aware_dro_repro.matrix import expand_tasks
from loss_aware_dro_repro.residuals import conic_residual_maximum
@pytest.fixture(scope="module")
def scientific_case():
task = next(task for task in expand_tasks(load_plan()) if task["task_id"] == "portfolio_gaussian_main/d001/r00/n010")
samples, _ = generate_dataset(task)
reference = gaussian_moments(samples)
bootstraps = bootstrap_gaussian_moments(samples, 20, task["seeds"]["bootstrap"])
L0 = np.eye(3)
epsilon, _ = initial_radius(reference, bootstraps, L0, 0.1)
return reference, bootstraps, L0, epsilon
def test_gaussian_portfolio_socp_residuals_and_decision(scientific_case):
reference, _, L0, epsilon = scientific_case
result = solve_gaussian_value(reference[0], reference[1], 0.05, epsilon, L0)
assert result["status"] == "optimal"
assert conic_residual_maximum(result["residuals"]) < 1e-8
assert abs(result["decision"].sum() - 1.0) < 1e-9
assert result["decision"].min() > -1e-9
expected_complementarity = abs(float(result["dual"] @ result["slack"]))
assert result["residuals"]["complementarity"] == pytest.approx(
expected_complementarity, rel=0.0, abs=1e-18
)
def test_clarabel_reduced_accuracy_cannot_exceed_frozen_residual_gate():
settings = clarabel_settings_contract()
assert load_plan()["numerics"]["solver_residual_max"] == 1e-7
assert settings["primary"]["tol_feas"] == 1e-10
assert settings["refinement"]["reduced_tol_feas"] <= 1e-8
assert settings["refinement_trigger_residual"] <= 5e-8
assert settings["refinement"]["iterative_refinement_max_iter"] == 50
assert settings["refinement"]["iterative_refinement_abstol"] == 1e-14
def test_exact_v3_failed_state_is_refined_without_weakening_gate(monkeypatch):
from loss_aware_dro_repro import conic as conic_module
plan = load_plan()
task = next(
task
for task in expand_tasks(plan)
if task["task_id"] == "portfolio_gaussian_highdim/d007/r06/n010"
)
samples, _ = generate_dataset(task)
reference = gaussian_moments(samples)
bootstraps = bootstrap_gaussian_moments(
samples,
task["hyperparameters"]["bootstrap_count"],
task["seeds"]["bootstrap"],
)
metric_factor = np.eye(samples.shape[1])
epsilon = float(
np.quantile(
[
gelbrich_value_gradient(reference, bootstrap, metric_factor)[0]
for bootstrap in bootstraps
],
1.0 - task["hyperparameters"]["coverage_beta"],
method=plan["numerics"]["bootstrap_quantile_method"],
)
)
scientific = {
"gamma": task["hyperparameters"]["cvar_gamma"],
"epsilon": epsilon,
"beta": task["hyperparameters"]["coverage_beta"],
"eta": task["hyperparameters"]["indicator_eta"],
"penalty_lambda": task["hyperparameters"]["coverage_penalty_lambda"],
"risk_mode": plan["numerics"]["gaussian_risk_estimands"]["primary"],
}
prior_solver_contract = clarabel_settings_contract()
prior_solver_contract["refinement_trigger_residual"] = 1.0
monkeypatch.setattr(
conic_module,
"clarabel_settings_contract",
lambda: prior_solver_contract,
)
for _ in range(179):
state = evaluate_outer_state(
reference, bootstraps, metric_factor, **scientific
)
assert state["lower"]["accuracy_refinement"]["applied"] is False
gradient = np.clip(
np.tril(state["total_gradient"]),
task["hyperparameters"]["gradient_clip"][0],
task["hyperparameters"]["gradient_clip"][1],
)
metric_factor, _ = stabilize_metric_factor(
np.tril(
metric_factor
- task["hyperparameters"]["learning_rate"] * gradient
),
task["hyperparameters"]["metric_eigenvalue_clip"][0],
task["hyperparameters"]["metric_eigenvalue_clip"][1],
)
monkeypatch.setattr(
conic_module,
"clarabel_settings_contract",
clarabel_settings_contract,
)
result = evaluate_outer_state(
reference, bootstraps, metric_factor, **scientific
)["lower"]
refinement = result["accuracy_refinement"]
assert refinement["applied"] is True
assert refinement["primary_status"] == "AlmostSolved"
assert refinement["primary_residual_maximum"] == pytest.approx(
2.0723689746484307e-07, rel=1e-12
)
assert result["raw_status"] == "Solved"
assert conic_residual_maximum(result["residuals"]) < 1e-9
def test_accuracy_refinement_fails_closed_if_second_solve_still_exceeds_trigger(
scientific_case, monkeypatch
):
from loss_aware_dro_repro import conic as module
reference, _, metric_factor, epsilon = scientific_case
original = module._solution_payload
calls = 0
def force_bad_residual(program, solution):
nonlocal calls
calls += 1
payload = original(program, solution)
payload["residuals"]["solver_native_primal"] = 1e-4
return payload
monkeypatch.setattr(module, "_solution_payload", force_bad_residual)
with pytest.raises(ContractError, match="accuracy refinement failed"):
solve_gaussian_value(
reference[0], reference[1], 0.05, epsilon, metric_factor
)
assert calls == 2
def test_conic_dual_envelope_matches_resolved_finite_difference(scientific_case):
reference, _, L0, epsilon = scientific_case
result = solve_gaussian_value(reference[0], reference[1], 0.05, epsilon, L0)
numerical = lower_triangular_central_difference(lambda L: solve_gaussian_value(reference[0], reference[1], 0.05, epsilon, L)["objective"], L0, 1e-5)
assert relative_gradient_error(result["value_gradient"], numerical) < 1e-6
def test_gelbrich_analytic_gradient_matches_finite_difference(scientific_case):
reference, bootstraps, L0, _ = scientific_case
_, analytic = gelbrich_value_gradient(reference, bootstraps[0], L0)
numerical = lower_triangular_central_difference(lambda L: gelbrich_value_gradient(reference, bootstraps[0], L)[0], L0, 1e-5)
assert relative_gradient_error(np.tril(analytic), numerical) < 1e-4
def test_gelbrich_exact_match_returns_zero_conservative_subgradient(scientific_case):
reference, _, L0, _ = scientific_case
value, gradient = gelbrich_value_gradient(reference, reference, L0)
assert value == 0.0
assert np.array_equal(gradient, np.zeros_like(L0))
def test_gelbrich_keeps_tiny_but_distinct_mean_shift():
first = (np.zeros(2), np.eye(2))
second = (np.array([1e-9, 0.0]), np.eye(2))
value, gradient = gelbrich_value_gradient(first, second, np.eye(2))
assert value == pytest.approx(1e-9, rel=1e-12)
assert gradient[0, 0] == pytest.approx(1e-9, rel=1e-12)
def test_active_coverage_penalty_gradient_matches_finite_difference(scientific_case):
reference, bootstraps, L0, epsilon = scientific_case
active_L = 1.2 * L0
kwargs = {"gamma": 0.05, "epsilon": epsilon, "beta": 0.1, "eta": 100.0, "penalty_lambda": 10.0}
state = evaluate_outer_state(reference, bootstraps, active_L, **kwargs)
assert state["penalty"] > 0
numerical = lower_triangular_central_difference(lambda L: evaluate_outer_state(reference, bootstraps, L, **kwargs)["total_objective"], active_L, 1e-5)
assert relative_gradient_error(state["total_gradient"], numerical) < 1e-4
def test_metric_projection_preserves_declared_eigenvalue_bounds():
proposed = np.array([[1e-8, 0.0], [2.0, 1e4]])
factor, receipt = stabilize_metric_factor(proposed, 1e-6, 1e6)
eigenvalues = np.linalg.eigvalsh(factor @ factor.T)
assert eigenvalues.min() >= 1e-6 - 1e-9
assert eigenvalues.max() <= 1e6 + 1e-6
assert receipt["metric_projection_fro"] > 0
def test_bilevel_steps_gates_on_paper_total_phi_and_keeps_released_diagnostic(monkeypatch):
from loss_aware_dro_repro import bilevel as module
lower_objectives = iter([10.0, 9.0, 8.0])
total_objectives = iter([10.0, 10.1, 9.0])
def fake_state(*args, **kwargs):
return {
"lower": {"objective": next(lower_objectives)},
"total_objective": next(total_objectives),
"total_gradient": np.zeros((1, 1)),
}
monkeypatch.setattr(module, "evaluate_outer_state", fake_state)
monkeypatch.setattr(module, "stabilize_metric_factor", lambda proposed, minimum, maximum: (proposed, {}))
config = {
"max_outer_iterations": 10,
"cvar_gamma": 0.05,
"epsilon": 1.0,
"coverage_beta": 0.1,
"indicator_eta": 100.0,
"coverage_penalty_lambda": 10.0,
"risk_coefficient_mode": "gaussian_exact",
"gradient_clip": [-1000.0, 1000.0],
"learning_rate": 1e-4,
"metric_eigenvalue_clip": [1e-6, 1e6],
"relative_objective_improvement_tolerance": 1e-6,
}
trace = bilevel_steps((np.zeros(1), np.eye(1)), [], np.eye(1), config)
assert len(trace) == 2
stopping = trace[-1]["stopping"]
assert stopping["paper_total_phi_literal"]["signed_relative_improvement"] < 0
assert stopping["paper_total_phi_literal"]["stop_trigger_caused_by_worsening"] is True
assert stopping["released_lower_objective_abs_denominator"]["signed_relative_improvement"] > 0
assert stopping["released_lower_objective_abs_denominator"]["stop_trigger_observed"] is False
assert stopping["reason"] == "paper_total_penalized_objective_worsened"
def test_stopping_contract_flags_negative_literal_denominator_and_does_not_use_safety_gate():
stopping = stopping_diagnostics(
previous_total_objective=-10.0,
current_total_objective=-11.0,
previous_lower_objective=5.0,
current_lower_objective=5.1,
tolerance=1e-6,
at_iteration_cap=False,
)
literal = stopping["paper_total_phi_literal"]
safety = stopping["paper_total_phi_abs_denominator"]
assert literal["signed_relative_improvement"] == pytest.approx(-0.1)
assert literal["denominator_sign_inversion_risk"] is True
assert literal["stop_triggered"] is True
assert safety["signed_relative_improvement"] == pytest.approx(0.1)
assert safety["stop_trigger_observed"] is False
assert stopping["reason"] == "paper_total_penalized_objective_improvement_below_tolerance"
def test_stopping_contract_zero_denominator_is_explicitly_undefined_and_cannot_trigger():
stopping = stopping_diagnostics(
previous_total_objective=0.0,
current_total_objective=1.0,
previous_lower_objective=0.0,
current_lower_objective=1.0,
tolerance=1e-6,
at_iteration_cap=False,
)
for name in (
"paper_total_phi_literal",
"paper_total_phi_abs_denominator",
"released_lower_objective_abs_denominator",
):
assert stopping[name]["signed_relative_improvement"] is None
assert stopping[name]["denominator_defined"] is False
assert stopping["reason"] is None

Xet Storage Details

Size:
11.9 kB
·
Xet hash:
eee97d44fd19a20666d0f9128ea3e4cad7b24a5f9c260b75774a89f93be01691

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.