File size: 1,210 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
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))