| 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") |
|
|