repro-raco-bundle / src /synthetic.py
junwatu's picture
Add reproduction bundle for RACO paper
1c1d9aa verified
Raw
History Blame Contribute Delete
2.81 kB
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