| 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) |
|
|