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) # Compute gp norms dev = g1.device dtype = g1.dtype # CAGrad (unclipped) 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)) # RACO (clipped) 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))