PIVOT / src /pivot /evaluation /rewards.py
pranamanam's picture
Upload 176 files
6fa9282 verified
Raw History Blame Contribute Delete
3.27 kB
"""Differentiable rewards for a specified target population."""
from __future__ import annotations
import torch
def rbf_mmd2(x: torch.Tensor, y: torch.Tensor, gamma: float) -> torch.Tensor:
"""Biased RBF MMD squared for (n,d) and (m,d) populations."""
k = lambda a, b: torch.exp(-gamma * torch.cdist(a, b).square())
return k(x, x).mean() + k(y, y).mean() - 2 * k(x, y).mean()
class Reward:
"""Per-cell or population reward; call returns (n,) for a population (n,d).
Centroid reward averages squared cell distances, including population
spread. Cosine reward averages per-cell effect cosines relative to a fixed
control centroid. MMD uses a training-derived fixed bandwidth.
"""
def __init__(
self,
kind="centroid",
target_c=None,
target_sample=None,
device="cpu",
control_ref=None,
gamma=None,
):
if kind not in ("centroid", "cosine", "mmd", "wasserstein", "nn_target"):
raise ValueError(f"Unsupported reward {kind}")
self.kind = kind
self.target_sample = (
torch.as_tensor(target_sample, dtype=torch.float32, device=device)
if target_sample is not None
else None
)
self.target_c = (
torch.as_tensor(target_c, dtype=torch.float32, device=device)
if target_c is not None
else self.target_sample.mean(0)
)
self.control_ref = (
torch.as_tensor(control_ref, dtype=torch.float32, device=device)
if control_ref is not None
else None
)
self.gamma = gamma
if kind == "cosine" and self.control_ref is None:
raise ValueError("Cosine requires a fixed control reference")
if kind == "mmd" and (gamma is None or gamma <= 0):
raise ValueError("MMD requires a positive fixed gamma")
if kind in ("mmd", "wasserstein", "nn_target") and self.target_sample is None:
raise ValueError("Population reward needs target samples")
def __call__(self, chat: torch.Tensor) -> torch.Tensor:
if self.kind == "centroid":
return -(chat - self.target_c).square().sum(-1)
if self.kind == "cosine":
p = chat - self.control_ref
t = self.target_c - self.control_ref
return (p * t).sum(-1) / (p.norm(dim=-1) * t.norm() + 1e-8)
if self.kind == "mmd":
return -rbf_mmd2(chat, self.target_sample, self.gamma).expand(len(chat))
if self.kind == "nn_target":
return -torch.cdist(chat, self.target_sample).square().min(1).values
# Fixed one-dimensional projections permit differentiable sorting. Quantile
# interpolation supports unequal source and target population sizes.
gen = torch.Generator(device=chat.device).manual_seed(0)
w = torch.randn(50, chat.shape[1], generator=gen, device=chat.device)
w = w / w.norm(dim=1, keepdim=True)
xp = chat @ w.T
yp = self.target_sample @ w.T
q = torch.linspace(0, 1, 256, device=chat.device)
cost = (
(torch.quantile(xp, q, dim=0) - torch.quantile(yp, q, dim=0)).abs().mean()
)
return -cost.expand(len(chat))