| import torch, sys, json, math |
| sys.path.insert(0, "..") |
| from src.moo import cagrad_update_k2, raco_update_k2, _grads_dot, _solve_lambda_k2 |
|
|
| torch.manual_seed(0) |
| w1, w2 = 0.8, 0.2 |
| c = 0.5 |
| 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) |
| p1_raw = float(lam_raw) |
| p2_raw = 1.0 - float(lam_raw) |
|
|
| p1_clipped = min(p1_raw, w1) |
| p2_clipped = min(p2_raw, w2) |
|
|
| |
| dev = g1.device |
| dtype = g1.dtype |
|
|
| |
| p1_t_u = torch.tensor(p1_raw, device=dev, dtype=dtype) |
| p2_t_u = torch.tensor(p2_raw, device=dev, dtype=dtype) |
| gp_norm_sq_u = (p1_t_u * p1_t_u) * H11 + (2.0 * p1_t_u * p2_t_u) * H12 + (p2_t_u * p2_t_u) * H22 |
| gp_norm_u = torch.sqrt(torch.clamp(gp_norm_sq_u, min=0.0)) |
|
|
| |
| p1_t_c = torch.tensor(p1_clipped, device=dev, dtype=dtype) |
| p2_t_c = torch.tensor(p2_clipped, device=dev, dtype=dtype) |
| gp_norm_sq_c = (p1_t_c * p1_t_c) * H11 + (2.0 * p1_t_c * p2_t_c) * H12 + (p2_t_c * p2_t_c) * H22 |
| gp_norm_c = torch.sqrt(torch.clamp(gp_norm_sq_c, min=0.0)) |
|
|
| a1_u = w1 + (s_val / float(gp_norm_u)) * p1_raw if float(gp_norm_u) > 0 else w1 |
| a2_u = w2 + (s_val / float(gp_norm_u)) * p2_raw if float(gp_norm_u) > 0 else w2 |
| a1_c = w1 + (s_val / float(gp_norm_c)) * p1_clipped if float(gp_norm_c) > 0 else w1 |
| a2_c = w2 + (s_val / float(gp_norm_c)) * p2_clipped if float(gp_norm_c) > 0 else w2 |
|
|
| g_update_u = g1 * a1_u + g2 * a2_u |
| g_update_c = g1 * a1_c + g2 * a2_c |
|
|
| result = { |
| "w1": w1, "w2": w2, |
| "p1_raw": p1_raw, "p2_raw": p2_raw, |
| "p1_clipped": p1_clipped, "p2_clipped": p2_clipped, |
| "gp_norm_cagrad": float(gp_norm_u), |
| "gp_norm_raco": float(gp_norm_c), |
| "coeffs_cagrad": {"a1": a1_u, "a2": a2_u}, |
| "coeffs_raco": {"a1": a1_c, "a2": a2_c}, |
| "correction_ratio_cagrad": {"obj1": a1_u - w1, "obj2": a2_u - w2}, |
| "correction_ratio_raco": {"obj1": a1_c - w1, "obj2": a2_c - w2}, |
| "raco_correction_more_balanced": (a1_u - w1) / (a2_u - w2 + 1e-10) < (a1_c - w1) / (a2_c - w2 + 1e-10) if (a2_u - w2) != 0 and (a2_c - w2) != 0 else None, |
| } |
| print(json.dumps(result, indent=2)) |
|
|