import torch import numpy as np from .moo import cagrad_update_k2, raco_update_k2 class ConflictingQuadratics: def __init__(self, dim=100, conflict_angle=np.pi * 0.6, seed=42): rng = np.random.RandomState(seed) self.dim = dim self.A1 = rng.randn(dim, dim).astype(np.float32) self.A1 = self.A1.T @ self.A1 / dim self.b1 = rng.randn(dim).astype(np.float32) * 0.1 self.A2 = rng.randn(dim, dim).astype(np.float32) self.A2 = self.A2.T @ self.A2 / dim self.b2 = rng.randn(dim).astype(np.float32) * 0.1 v = rng.randn(dim).astype(np.float32) v = v / np.linalg.norm(v) u = rng.randn(dim).astype(np.float32) u = u - np.dot(u, v) * v u = u / np.linalg.norm(u) g1_0 = v g2_0 = v * np.cos(conflict_angle) + u * np.sin(conflict_angle) scale = np.linalg.norm(self.A1 @ self.b1 + self.b1) self.b1 = self.b1 - g1_0 * scale self.b2 = self.b2 - g2_0 * scale def run_synthetic_comparison(dim=100, conflict_angle=np.pi*0.6, steps=200, lr=0.01, seed=42): import math prob = ConflictingQuadratics(dim=dim, conflict_angle=conflict_angle, seed=seed) A1_t = torch.from_numpy(prob.A1).float() b1_t = torch.from_numpy(prob.b1).float() A2_t = torch.from_numpy(prob.A2).float() b2_t = torch.from_numpy(prob.b2).float() theta0 = torch.from_numpy(np.random.RandomState(seed + 1).randn(dim).astype(np.float32) * 0.5) methods = {} for name, use_clip, w1, w2 in [ ("DPO_LW", False, 0.7, 0.3), ("CAGrad", False, 0.7, 0.3), ("RACO", True, 0.7, 0.3), ]: theta = theta0.clone().requires_grad_(True) losses = {0: [], 1: []} opt = torch.optim.SGD([theta], lr=lr) for t in range(steps): opt.zero_grad() l1 = 0.5 * (theta @ A1_t @ theta) + b1_t @ theta l2 = 0.5 * (theta @ A2_t @ theta) + b2_t @ theta g1 = torch.autograd.grad(l1, theta, retain_graph=True)[0] g2 = torch.autograd.grad(l2, theta, retain_graph=True)[0] g1_list = [g1] g2_list = [g2] if name == "DPO_LW": g_update = g1 * w1 + g2 * w2 elif name == "CAGrad": gl = cagrad_update_k2(g1_list, g2_list, w1, w2, c=0.4) g_update = gl[0] elif name == "RACO": gl = raco_update_k2(g1_list, g2_list, w1, w2, c=0.4) g_update = gl[0] theta.grad = g_update.detach() if hasattr(g_update, 'detach') else g_update opt.step() losses[0].append(float(l1.detach())) losses[1].append(float(l2.detach())) methods[name] = {"theta": theta.detach().numpy().copy(), "losses": losses} return methods