junwatu's picture
Add reproduction bundle for RACO paper
1c1d9aa verified
Raw
History Blame Contribute Delete
3.59 kB
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)