File size: 3,268 Bytes
6fa9282
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
"""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))