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))