Download src/pivot/evaluation/rewards.py from ChatterjeeLab/PIVOT: direct link, hf CLI and curl.
- Browser
- Download file 3.27 kB
-
https://huggingface.co/ChatterjeeLab/PIVOT/resolve/main/src/pivot/evaluation/rewards.py
- Command line
-
hf download hf://ChatterjeeLab/PIVOT/src/pivot/evaluation/rewards.py
-
curl -L -o rewards.py https://huggingface.co/ChatterjeeLab/PIVOT/resolve/main/src/pivot/evaluation/rewards.py
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)) | |