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)