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