| import torch |
| import pytest |
| from optimal_transport.ot_fc_sep_map import FCOTSeparable |
| from models import FiniteSeparableModel |
|
|
|
|
| |
| |
| |
|
|
| def quad_kernel(x, y): |
| return (x - y) ** 2 |
|
|
|
|
| |
| |
| |
|
|
| @pytest.fixture |
| def small_model(): |
| model = FiniteSeparableModel( |
| kernel=quad_kernel, |
| num_dims=1, |
| radius=2.0, |
| y_accuracy=0.5, |
| x_accuracy=0.5, |
| mode="concave", |
| temp=5.0, |
| epsilon=1e-6, |
| cache_gradients=False, |
| ) |
| return model |
|
|
|
|
| @pytest.fixture |
| def small_solver(small_model): |
| solver = FCOTSeparable( |
| input_dim=1, |
| model=small_model, |
| inverse_kx=lambda x, p: x - 0.5 * p, |
| outer_lr=1e-2, |
| warmup_lr=1e-1, |
| full_refresh_every=5, |
| reactivate_every=3, |
| reactivate_eps=1e-3, |
| ) |
| return solver |
|
|
|
|
| |
| |
| |
|
|
| def test_refresh_updates_intercepts(small_model): |
| |
| torch.manual_seed(0) |
| with torch.no_grad(): |
| small_model.intercepts.copy_(torch.randn_like(small_model.intercepts)) |
|
|
| old = small_model.intercepts.clone() |
|
|
| changed = small_model.refresh_intercepts_via_transform() |
| new = small_model.intercepts |
|
|
| |
| |
| assert isinstance(changed, int) |
| assert new.shape == old.shape |
|
|
|
|
| |
| |
| |
|
|
| def test_momentum_reset(small_solver): |
| solver = small_solver |
| opt = solver.optimizer |
|
|
| p = solver.model.intercepts |
|
|
| |
| for group in opt.param_groups: |
| for param in group["params"]: |
| if param is p: |
| state = opt.state.setdefault(param, {}) |
| state["momentum_buffer"] = torch.ones_like(p) * 3.14 |
|
|
| |
| solver._reset_intercept_momentum() |
|
|
| |
| state = opt.state[p] |
| assert len(state) == 0, "Momentum/state was not fully cleared." |
|
|
|
|
| |
| |
| |
|
|
| def test_jitter_changes_intercepts(small_solver): |
| model = small_solver.model |
| before = model.intercepts.clone() |
| small_solver._jitter_intercepts(strength=1e-2) |
| after = model.intercepts |
|
|
| assert not torch.allclose(before, after), "Jitter must modify intercepts." |
| diff = (after - before).abs().mean().item() |
| assert diff > 0, "Jitter must have nonzero magnitude." |
|
|
|
|
| |
| |
| |
|
|
| def test_full_refresh_triggers_all(small_solver): |
| solver = small_solver |
| model = solver.model |
|
|
| |
| calls = {"refresh": 0, "reset": 0, "jitter": 0} |
|
|
| orig_refresh = model.refresh_intercepts_via_transform |
| def counted_refresh(): |
| calls["refresh"] += 1 |
| return orig_refresh() |
|
|
| solver.model.refresh_intercepts_via_transform = counted_refresh |
|
|
| orig_reset = solver._reset_intercept_momentum |
| def counted_reset(): |
| calls["reset"] += 1 |
| return orig_reset() |
|
|
| solver._reset_intercept_momentum = counted_reset |
|
|
| orig_jitter = solver._jitter_intercepts |
| def counted_jitter(strength=None): |
| calls["jitter"] += 1 |
| return orig_jitter(strength) |
|
|
| solver._jitter_intercepts = counted_jitter |
|
|
| |
| solver._step_count = solver._full_refresh_every |
| solver._maybe_full_refresh() |
|
|
| assert calls["refresh"] == 1, "Refresh should be called exactly once." |
| assert calls["reset"] == 1, "Momentum reset should be called exactly once." |
| assert calls["jitter"] == 1, "Jitter should be called exactly once." |
|
|
|
|
| |
| |
| |
|
|
| def test_forward_shapes(small_model): |
| X = torch.tensor([[0.2], [-0.5], [1.0]]) |
| choice, fx = small_model.forward(X, selection_mode="hard") |
|
|
| assert choice.shape == (3, 1) |
| assert fx.shape == (3,) |
| assert not fx.isnan().any() |
|
|
|
|
| def test_ste_gradient_correctness(small_model): |
| |
| X = torch.tensor([[0.1], [0.5]], requires_grad=True) |
| _, fx = small_model.forward(X, selection_mode="hard") |
| loss = fx.sum() |
| loss.backward() |
|
|
| assert X.grad is not None, "Gradient must flow through STE wrapper." |
| assert not torch.isnan(X.grad).any() |
| |
| assert X.grad.abs().max() > 0 |
|
|
|
|
| |
| |
| |
|
|
| def test_transform_consistency(small_model): |
| Z = torch.linspace(-1, 1, 5).reshape(-1, 1) |
| Xopt, vals, conv = small_model.inf_transform(Z) |
|
|
| assert Xopt.shape == Z.shape, "Transform output must match input shape." |
| assert vals.shape == (5,), "Value shape mismatch." |
| assert conv is True, "Separable transform must always converge." |
| assert not torch.isnan(vals).any() |
|
|
|
|
| |
| |
| |
|
|
| def test_training_makes_progress(): |
| """ |
| Non-monotone-safe training test. |
| |
| Checks that: |
| - objective stays finite, |
| - gradients stay reasonable, |
| - progress occurs over a window, |
| without assuming monotonic ascent (refresh/jitter/STE break monotonicity). |
| """ |
|
|
| import math |
| import torch |
| from optimal_transport.ot_fc_sep_map import FCOTSeparable |
|
|
| torch.manual_seed(0) |
|
|
| |
| X = torch.randn(200, 1) * 0.5 |
| Y = torch.randn(200, 1) + 1.0 |
|
|
| |
| solver = FCOTSeparable.initialize_right_architecture( |
| dim=1, |
| radius=3.0, |
| n_params=40, |
| x_accuracy=0.1, |
| kernel_1d=lambda x, y: (x - y) ** 2, |
| inverse_kx=lambda x, p: x - 0.5 * p, |
| outer_lr=5e-3, |
| warmup_lr=1e-2, |
| warmup_grad_threshold=0.1, |
| warmup_max_steps=20, |
| full_refresh_every=10, |
| reactivate_every=15, |
| ) |
|
|
| history = [] |
| for _ in range(50): |
| res = solver.step(X, Y) |
| obj = res["dual"] |
| grad_norm = res["grad_norm"] |
|
|
| |
| assert not math.isnan(obj) |
| assert not math.isinf(obj) |
| assert grad_norm < 1e6 |
|
|
| history.append(obj) |
|
|
| |
| early_best = max(history[:5]) |
| late_best = max(history[-10:]) |
|
|
| assert late_best > early_best, ( |
| f"OT objective failed to improve over time: " |
| f"early={early_best:.4f}, late={late_best:.4f}" |
| ) |
|
|
| |
| |
| |
|
|
| def test_transport_map(small_solver): |
| solver = small_solver |
| X = torch.tensor([[-0.5], [0.0], [0.5]], dtype=torch.float32) |
|
|
| |
| X_requires_grad = X.to(solver.device).requires_grad_(True) |
| _, u_X = solver.model.forward(X_requires_grad, selection_mode="hard") |
| grad_u = torch.autograd.grad(u_X.sum(), X_requires_grad, create_graph=False)[0] |
| Ypred = solver.inverse_kx(X_requires_grad.detach(), grad_u.detach()) |
|
|
| assert Ypred.shape == X.shape |
| assert not torch.isnan(Ypred).any() |
| |
| assert Ypred.abs().max() < 10, "Transport values look unreasonable." |
|
|