File size: 1,642 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
import torch, sys, json, math
sys.path.insert(0, "..")
from src.moo import cagrad_update_k2, raco_update_k2, _grads_dot, _solve_lambda_k2

seeds = [0, 1, 7, 13, 42, 99, 123, 256]
w1, w2 = 0.8, 0.2
results = []

for seed in seeds:
    for c in [0.3, 0.5, 0.7]:
        torch.manual_seed(seed)
        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)
        p2_raw = 1.0 - float(lam_raw)

        if p2_raw > w2:
            cagrad_g = cagrad_update_k2([g1], [g2], w1, w2, c=c)
            raco_g = raco_update_k2([g1], [g2], w1, w2, c=c)

            ca_alignment_g2 = float((cagrad_g[0] @ g2).item())
            ra_alignment_g2 = float((raco_g[0] @ g2).item())

            results.append({
                "seed": seed, "c": c,
                "p2_raw": p2_raw, "w2": w2,
                "overcorrection": p2_raw > w2,
                "cagrad_g2_improvement": ca_alignment_g2,
                "raco_g2_improvement": ra_alignment_g2,
                "raco_reduces_overcorrection": abs(ra_alignment_g2) < abs(ca_alignment_g2),
            })

for r in results[:5]:
    print(json.dumps(r, indent=2))

print(f"\nFound {len(results)} cases with overcorrection out of {len(seeds) * 3} attempts")