File size: 3,594 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
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
import torch
import torch.nn.functional as F
import numpy as np

def _get_grad(g):
    if isinstance(g, tuple):
        return g[0]
    return g

def _grads_dot(ga, gb):
    dot = 0.0
    for a, b in zip(ga, gb):
        a = _get_grad(a)
        b = _get_grad(b)
        if a is not None and b is not None:
            dot = dot + (a.flatten() @ b.flatten())
    return dot

def compute_dpo_loss(beta, delta, sign):
    return -F.logsigmoid(beta * sign * delta).mean()

def _solve_lambda_k2(b1, b2, H11, H12, H22, s):
    d = (b2 - b1) + (H11 - 2 * H12 + H22) * 0.5
    eps = 1e-10
    if abs(d) > eps:
        lam = 0.5 * (1.0 + (b1 - b2) / d)
        lam = max(0.0, min(1.0, lam))
    else:
        lam = 0.5
    return lam, {"case": "closed_form"}

def _build_update(g1, g2, w1, w2, p1, p2, s_val):
    g1_0 = _get_grad(g1[0])
    g2_0 = _get_grad(g2[0])
    dev = g1_0.device if g1_0 is not None else g2_0.device
    dtype = g1_0.dtype if g1_0 is not None else g2_0.dtype

    p1_t = torch.tensor(p1, device=dev, dtype=dtype)
    p2_t = torch.tensor(p2, device=dev, dtype=dtype)

    gp_norm_sq = (p1_t * p1_t) * _grads_dot(g1, g1) \
                 + (2.0 * p1_t * p2_t) * _grads_dot(g1, g2) \
                 + (p2_t * p2_t) * _grads_dot(g2, g2)
    gp_norm = torch.sqrt(torch.clamp(gp_norm_sq, min=0.0))

    if float(gp_norm) > 0.0:
        scale_p = float(s_val / gp_norm)
        a1 = w1 + scale_p * p1
        a2 = w2 + scale_p * p2
    else:
        a1, a2 = w1, w2

    result = []
    for g1_i, g2_i in zip(g1, g2):
        g1_i = _get_grad(g1_i)
        g2_i = _get_grad(g2_i)
        if g1_i is None and g2_i is None:
            result.append(None)
        elif g1_i is None:
            result.append(g2_i.detach() * a2)
        elif g2_i is None:
            result.append(g1_i.detach() * a1)
        else:
            result.append(g1_i.detach() * a1 + g2_i.detach() * a2)
    return result

def cagrad_update_k2(g1, g2, w1, w2, c=0.4):
    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))
    if float(g0_norm) == 0.0 or c == 0.0:
        return [(_get_grad(g1_i) * w1 + _get_grad(g2_i) * w2) if _get_grad(g1_i) is not None and _get_grad(g2_i) is not None else (_get_grad(g1_i) * w1 if _get_grad(g1_i) is not None else _get_grad(g2_i) * w2) for g1_i, g2_i in zip(g1, g2)]

    b1 = (w1 * H11) + (w2 * H12)
    b2 = (w1 * H12) + (w2 * H22)
    s_val = float(g0_norm) * c

    lam, _ = _solve_lambda_k2(b1, b2, H11, H12, H22, s_val)
    p1 = float(lam)
    p2 = 1.0 - float(lam)

    return _build_update(g1, g2, w1, w2, p1, p2, s_val)

def raco_update_k2(g1, g2, w1, w2, c=0.4):
    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))
    if float(g0_norm) == 0.0 or c == 0.0:
        return [(_get_grad(g1_i) * w1 + _get_grad(g2_i) * w2) if _get_grad(g1_i) is not None and _get_grad(g2_i) is not None else (_get_grad(g1_i) * w1 if _get_grad(g1_i) is not None else _get_grad(g2_i) * w2) for g1_i, g2_i in zip(g1, g2)]

    b1 = (w1 * H11) + (w2 * H12)
    b2 = (w1 * H12) + (w2 * H22)
    s_val = float(g0_norm) * c

    lam, _ = _solve_lambda_k2(b1, b2, H11, H12, H22, s_val)
    p1_raw = float(lam)
    p2_raw = 1.0 - float(lam)

    p1 = min(p1_raw, w1)
    p2 = min(p2_raw, w2)

    return _build_update(g1, g2, w1, w2, p1, p2, s_val)