| import torch, sys, json |
| sys.path.insert(0, "..") |
| from src.moo import cagrad_update_k2, raco_update_k2, _grads_dot, _solve_lambda_k2 |
|
|
| torch.manual_seed(42) |
| w1, w2 = 0.8, 0.2 |
|
|
| g1 = torch.randn(100) * 3.0 |
| g2 = -g1 * 0.9 + torch.randn(100) * 0.3 |
|
|
| 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) |
| c = 0.5 |
| s_val = float(g0_norm) * c |
|
|
| lam_raw, dbg = _solve_lambda_k2(b1, b2, H11, H12, H22, s_val) |
|
|
| p1_raw = float(lam_raw) |
| p2_raw = 1.0 - float(lam_raw) |
|
|
| p1_clipped = min(p1_raw, w1) |
| p2_clipped = min(p2_raw, w2) |
|
|
| data = { |
| "H11": float(H11), |
| "H12": float(H12), |
| "H22": float(H22), |
| "g0_norm": float(g0_norm), |
| "b1": float(b1), |
| "b2": float(b2), |
| "s_val": s_val, |
| "lam_raw": lam_raw, |
| "p1_raw": p1_raw, |
| "p2_raw": p2_raw, |
| "w1": w1, |
| "w2": w2, |
| "p1_clipped": p1_clipped, |
| "p2_clipped": p2_clipped, |
| "p1_was_clipped": p1_raw > w1, |
| "p2_was_clipped": p2_raw > w2, |
| "any_clipping": p1_raw > w1 or p2_raw > w2, |
| } |
| print(json.dumps(data, indent=2)) |
|
|