99 lines
3.0 KiB
Python
99 lines
3.0 KiB
Python
"""PKPO reward transformations — Listing 1 of the paper (arXiv:2505.15201), verbatim.
|
|
|
|
sloo_minus_one is the s^(loo-1) estimator of Eq. (33)/(34): transformed rewards whose
|
|
sum estimates pass@k / max_g@k, with a k-1-subset LOO baseline. Requires n >= k+1.
|
|
At k=1 the paper uses untransformed rewards with mean centering (its k-1 baseline is
|
|
undefined); transform_rewards() below handles that case.
|
|
"""
|
|
from typing import Callable
|
|
|
|
import numpy as np
|
|
|
|
|
|
def _m_normed(N: int, K: int, i: int, j: int) -> float:
|
|
if i == j and i >= K - 1:
|
|
return (
|
|
K / (N - K + 1)
|
|
* np.prod(np.arange(i - K + 2, i + 1) / np.arange(N - K + 2, N + 1))
|
|
)
|
|
elif j > i and j >= K - 1 and K >= 2:
|
|
return (
|
|
K / (N - K + 1) * (K - 1) / N
|
|
* np.prod(np.arange(j - K + 2, j) / np.arange(N - K + 2, N))
|
|
)
|
|
return 0
|
|
|
|
|
|
def _m_diagonal(N: int, K: int) -> np.ndarray:
|
|
return np.array([_m_normed(N, K, i, i) for i in range(N)])
|
|
|
|
|
|
def rho(g: np.ndarray, K: int) -> float:
|
|
"""See Equation (12)."""
|
|
return (np.sort(g) * _m_diagonal(len(g), K)).sum()
|
|
|
|
|
|
def _delta(N: int, K: int, i: int) -> float:
|
|
return _m_normed(N, K, i, i + 1) - _m_normed(N, K, i + 1, i + 1)
|
|
|
|
|
|
def _deltas(N: int, K: int) -> np.ndarray:
|
|
return np.array([_delta(N - 1, K, i) for i in range(N - 2)])
|
|
|
|
|
|
def _sorted_apply(func: Callable) -> Callable:
|
|
def inner(x: np.ndarray, *args, **kwargs) -> np.ndarray:
|
|
i_sort = np.argsort(x)
|
|
func_x = np.zeros_like(x)
|
|
func_x[i_sort] = func(x[i_sort], *args, **kwargs)
|
|
return func_x
|
|
return inner
|
|
|
|
|
|
@_sorted_apply
|
|
def s(g: np.ndarray, K: int):
|
|
"""See Equation (19)."""
|
|
N = len(g)
|
|
c = g * _m_diagonal(N, K)
|
|
c[:(N - 1)] += g[1:] * _deltas(N + 1, K)
|
|
return np.cumsum(c[::-1])[::-1]
|
|
|
|
|
|
@_sorted_apply
|
|
def _b(g: np.ndarray, K: int) -> np.ndarray:
|
|
N = len(g)
|
|
w = (_m_diagonal(N - 1, K) * np.arange(1, N)).astype(float)
|
|
w[1:] += _deltas(N, K) * np.arange(1, N - 1)
|
|
c1 = np.array([(w * g[1:]).sum()])
|
|
c2 = (g[:-1] - g[1:]) * w
|
|
return np.cumsum(np.concatenate((c1, c2)))
|
|
|
|
|
|
def sloo(g: np.ndarray, K: int) -> np.ndarray:
|
|
"""See Equation (29)."""
|
|
return s(g, K) - _b(g, K) / (len(g) - 1)
|
|
|
|
|
|
def sloo_minus_one(g: np.ndarray, K: int) -> np.ndarray:
|
|
"""See Equation (33)."""
|
|
return s(g, K) - _b(g, K - 1) * K / (K - 1) / len(g)
|
|
|
|
|
|
def transform_rewards(g: np.ndarray, K: int) -> np.ndarray:
|
|
"""PKPO advantages for one group of n rollouts of the same prompt.
|
|
|
|
k >= 2: sloo_minus_one exactly as in Listing 1.
|
|
k == 1: s(g, 1) (= g/n, no transformation) with group-mean centering, which is
|
|
what the paper uses for its k_opt=1 runs ("without which the training diverges").
|
|
Scaled by n so advantage magnitude is O(reward) at every k (constant across
|
|
stages; equivalent to a learning-rate rescale).
|
|
"""
|
|
g = np.asarray(g, dtype=np.float64)
|
|
n = len(g)
|
|
if K >= 2:
|
|
out = sloo_minus_one(g, K)
|
|
else:
|
|
out = s(g, 1)
|
|
out = out - out.mean()
|
|
return out * n
|