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

w1, w2 = 0.8, 0.2
c = 0.4

for seed in range(200):
    torch.manual_seed(seed)
    g1 = torch.randn(100)
    g2 = torch.randn(100)

    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)

    if 0.0 < p1_raw < 1.0 and p2_raw > w2:
        # Clipping changes the ratio p1:p2
        p1_clipped = min(p1_raw, w1)
        p2_clipped = min(p2_raw, w2)

        # Compute the actual update directions
        cagrad_g = cagrad_update_k2([g1], [g2], w1, w2, c=c)
        raco_g = raco_update_k2([g1], [g2], w1, w2, c=c)

        diff_norm = float((cagrad_g[0] - raco_g[0]).norm().item())
        if diff_norm > 1e-4:
            g0 = g1 * w1 + g2 * w2
            ca_dir = cagrad_g[0] / cagrad_g[0].norm()
            ra_dir = raco_g[0] / raco_g[0].norm()
            g0_dir = g0 / g0.norm()

            print(f"seed={seed}: p_raw=({p1_raw:.3f},{p2_raw:.3f}) "
                  f"p_clipped=({p1_clipped:.3f},{p2_clipped:.3f}) "
                  f"diff_norm={diff_norm:.4f}")
            print(f"  CAGrad direction deviates from g0 by {float((ca_dir - g0_dir).norm().item()):.4f}")
            print(f"  RACO direction deviates from g0 by {float((ra_dir - g0_dir).norm().item()):.4f}")
            break
else:
    print("No clipping case found in 200 attempts - this is expected since the geometry often makes p_extreme=0 or 1")