repro-raco-bundle / tests /test_moo.py
junwatu's picture
Add reproduction bundle for RACO paper
1c1d9aa verified
Raw
History Blame Contribute Delete
2.82 kB
import torch, sys, json
sys.path.insert(0, "..")
from src.moo import cagrad_update_k2, raco_update_k2, _grads_dot, _solve_lambda_k2
def test_clipping_triggers():
"""When CAGrad wants a large correction on the less-preferred objective,
RACO should clip it."""
w1, w2 = 0.8, 0.2
torch.manual_seed(42)
g1 = torch.randn(100) * 3.0
g2 = -g1 * 0.9 + torch.randn(100) * 0.3
g1_list = [g1]
g2_list = [g2]
cagrad_g = cagrad_update_k2(g1_list, g2_list, w1, w2, c=0.5)
raco_g = raco_update_k2(g1_list, g2_list, w1, w2, c=0.5)
cg = cagrad_g[0]
rg = raco_g[0]
alignment_ca_g1 = float((cg @ g1).item())
alignment_ca_g2 = float((cg @ g2).item())
alignment_ra_g1 = float((rg @ g1).item())
alignment_ra_g2 = float((rg @ g2).item())
result = {
"test": "Strong conflicting objectives, w1=0.8, w2=0.2",
"w1": w1, "w2": w2,
"g1_norm": float(g1.norm().item()),
"g2_norm": float(g2.norm().item()),
"cos_angle": float((g1 @ g2 / (g1.norm() * g2.norm())).item()),
"cagrad_alignment_g1": alignment_ca_g1,
"cagrad_alignment_g2": alignment_ca_g2,
"raco_alignment_g1": alignment_ra_g1,
"raco_alignment_g2": alignment_ra_g2,
"cagrad_more_g2_magnitude": abs(alignment_ca_g2) > abs(alignment_ra_g2),
"raco_better_respects_weights": alignment_ra_g2 > alignment_ca_g2,
}
return result
def test_theorem_convergence():
"""Test that CAGrad-Clip provides valid descent directions."""
results = {}
for c in [0.1, 0.3, 0.5, 0.7]:
w1 = 0.7
w2 = 0.3
torch.manual_seed(123)
g1 = torch.randn(50) * 2.0
g2 = -g1 * 0.7 + torch.randn(50) * 0.5
g1_list, g2_list = [g1], [g2]
try:
rg = raco_update_k2(g1_list, g2_list, w1, w2, c=c)
g_update = rg[0]
g0 = g1 * w1 + g2 * w2
descent_ok = float((g_update @ g_update).item()) > 0
results[f"c={c}"] = {
"valid": descent_ok,
"g0_norm": float(g0.norm().item()),
"g_update_norm": float(g_update.norm().item()),
}
except Exception as e:
results[f"c={c}"] = {"error": str(e)}
return results
if __name__ == "__main__":
all_tests = {}
all_tests["test_clipping"] = test_clipping_triggers()
all_tests["test_convergence"] = test_theorem_convergence()
print(json.dumps(all_tests, indent=2))
with open("/Users/equan_p/Developer/playground/ICML-2/repro_raco/outputs/tier1/test_results.json", "w") as f:
json.dump(all_tests, f, indent=2)
all_pass = all(
t.get("raco_better_respects_weights", True) or t.get("valid", True)
for t in all_tests.values()
)
print(f"\nAll tests passed: {all_pass}")