File size: 2,387 Bytes
1c1d9aa | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 | 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))
|