import torch, sys, json, math sys.path.insert(0, "..") from src.moo import cagrad_update_k2, raco_update_k2, _grads_dot, _solve_lambda_k2 seeds = [0, 1, 7, 13, 42, 99, 123, 256] w1, w2 = 0.8, 0.2 results = [] for seed in seeds: for c in [0.3, 0.5, 0.7]: torch.manual_seed(seed) g1 = torch.randn(100) * 3.0 g2 = torch.randn(100) * 2.0 H11 = _grads_dot([g1], [g1]) H12 = _grads_dot([g1], [g2]) H22 = _grads_dot([g2], [g2]) g0_norm_sq = (w1 * w1) * H11 + (2.0 * w1 * w2) * H12 + (w2 * w2) * H22 g0_norm = torch.sqrt(torch.clamp(g0_norm_sq, min=0.0)) b1 = (w1 * H11) + (w2 * H12) b2 = (w1 * H12) + (w2 * H22) s_val = float(g0_norm) * c lam_raw, _ = _solve_lambda_k2(b1, b2, H11, H12, H22, s_val) p2_raw = 1.0 - float(lam_raw) if p2_raw > w2: cagrad_g = cagrad_update_k2([g1], [g2], w1, w2, c=c) raco_g = raco_update_k2([g1], [g2], w1, w2, c=c) ca_alignment_g2 = float((cagrad_g[0] @ g2).item()) ra_alignment_g2 = float((raco_g[0] @ g2).item()) results.append({ "seed": seed, "c": c, "p2_raw": p2_raw, "w2": w2, "overcorrection": p2_raw > w2, "cagrad_g2_improvement": ca_alignment_g2, "raco_g2_improvement": ra_alignment_g2, "raco_reduces_overcorrection": abs(ra_alignment_g2) < abs(ca_alignment_g2), }) for r in results[:5]: print(json.dumps(r, indent=2)) print(f"\nFound {len(results)} cases with overcorrection out of {len(seeds) * 3} attempts")