| import torch, sys, json |
| sys.path.insert(0, "..") |
| from src.moo import cagrad_update_k2, raco_update_k2, _grads_dot, _solve_lambda_k2 |
|
|
| w1, w2 = 0.8, 0.2 |
| c = 0.4 |
|
|
| for seed in range(200): |
| torch.manual_seed(seed) |
| g1 = torch.randn(100) |
| g2 = torch.randn(100) |
|
|
| 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) |
| p1_raw = float(lam_raw) |
| p2_raw = 1.0 - float(lam_raw) |
|
|
| if 0.0 < p1_raw < 1.0 and p2_raw > w2: |
| |
| p1_clipped = min(p1_raw, w1) |
| p2_clipped = min(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) |
|
|
| diff_norm = float((cagrad_g[0] - raco_g[0]).norm().item()) |
| if diff_norm > 1e-4: |
| g0 = g1 * w1 + g2 * w2 |
| ca_dir = cagrad_g[0] / cagrad_g[0].norm() |
| ra_dir = raco_g[0] / raco_g[0].norm() |
| g0_dir = g0 / g0.norm() |
|
|
| print(f"seed={seed}: p_raw=({p1_raw:.3f},{p2_raw:.3f}) " |
| f"p_clipped=({p1_clipped:.3f},{p2_clipped:.3f}) " |
| f"diff_norm={diff_norm:.4f}") |
| print(f" CAGrad direction deviates from g0 by {float((ca_dir - g0_dir).norm().item()):.4f}") |
| print(f" RACO direction deviates from g0 by {float((ra_dir - g0_dir).norm().item()):.4f}") |
| break |
| else: |
| print("No clipping case found in 200 attempts - this is expected since the geometry often makes p_extreme=0 or 1") |
|
|