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