Skip to main content

Reference semantics, conformance testing, and pathology diagnostics for LLM policy-gradient post-training

Project description

polgrad

CI PyPI Python License

Reference semantics, conformance testing, and pathology diagnostics for LLM policy-gradient post-training. Pure functions over tensors: no trainers, no models, no GPU requirement, no dependency beyond PyTorch.

The problem

The RL post-training stack does not agree on what "the GRPO loss" is. The same batch of (logprobs, old_logprobs, advantages, response_mask) produces materially different losses and gradients depending on which framework computes it, and on which version of that framework:

  • TRL's GRPOTrainer aggregated per-sequence-then-batch (the GRPO paper equation) in v0.14–v0.15, then switched its default to global token normalization in v0.16.0 (trl#2881; see trl#2995 for the debate that switch sparked, resolved for KL logging only in trl#3004). Both defaults are still in the wild. Current TRL defaults to loss_type="dapo": token-mean over the global gradient-accumulated batch.
  • verl ships four aggregation modes; its Dr.GRPO-style mode divides by the current padded length where the Dr.GRPO paper's released code divides by a fixed generation budget (arXiv:2503.20783). polgrad's conformance suite measures the exact factor.
  • The k3 KL estimator, differentiated as a loss term, does not propagate the gradient of the KL it estimates. Its pathwise gradient is biased for the stated objective (arXiv:2512.21852, arXiv:2510.01555). k3-as-loss is a common configuration anyway.
  • Rollout engines report logprobs that differ from the trainer's recompute, silently making "on-policy" training off-policy; frameworks patched in truncated importance-sampling corrections (verl#2953).

These are semantic choices with measurable consequences, scattered across codebases as implicit defaults. polgrad makes each one an explicit, typed, tested object.

What disagreement looks like

One ragged batch of 4 sequences with lengths (2, 4, 6, 8), and the same tensors into every row of this table. examples/which_grpo_am_i_running.py prints it, measured, on every run:

configuration loss grad mass, L=2 seq grad mass, L=8 seq
seq-mean-token-mean (GRPO paper; TRL ≤ 0.15) 0.24038972 0.19568 0.29263
token-mean (verl; TRL ≥ 0.16) 0.51717388 0.07827 0.46821
token-sum-norm (Dr.GRPO, budget 8) 0.32323368 0.04892 0.29263
GSPO sequence-level ratio 0.15429390 0.19550 0.00000

A 2× loss difference and opposite length-weighting of gradients, from settings frameworks treat as interchangeable defaults. The same script cross-checks polgrad against the vendored verl and OpenRLHF implementations and prints the agreement: 0 or ~1e-10, with each residual explained.

Install

pip install polgrad

or, for the development version, pip install git+https://github.com/SaifPunjwani/polgrad.

Python ≥ 3.10, PyTorch ≥ 2.1, CPU is enough.

The loss algebra

Every named variant is a frozen config over ratio kind × surrogate × clip × aggregation × KL placement. One code path evaluates all of them:

import torch
from polgrad import Aggregation, ClipConfig, PolicyLossConfig, RatioKind, SurrogateKind, policy_loss

grpo_paper = PolicyLossConfig(
    ratio=RatioKind.TOKEN,
    surrogate=SurrogateKind.PG_CLIP,
    clip=ClipConfig(eps_low=0.2, eps_high=0.2),
    aggregation=Aggregation.SEQ_MEAN_TOKEN_MEAN,
)
trl_016_default = PolicyLossConfig(
    ratio=RatioKind.TOKEN,
    surrogate=SurrogateKind.PG_CLIP,
    clip=ClipConfig(eps_low=0.2, eps_high=0.2),
    aggregation=Aggregation.TOKEN_MEAN,
)

B, T = 4, 8
g = torch.Generator().manual_seed(0)
logprobs = -torch.rand((B, T), generator=g, dtype=torch.float64)
old_logprobs = logprobs.detach() + 0.1 * torch.randn((B, T), generator=g, dtype=torch.float64)
advantages = torch.randn((B,), generator=g, dtype=torch.float64)
mask = torch.arange(T).expand(B, T) < torch.tensor([[2], [4], [6], [8]])

for cfg in (grpo_paper, trl_016_default):
    out = policy_loss(cfg, logprobs=logprobs.clone().requires_grad_(), old_logprobs=old_logprobs,
                      advantages=advantages, response_mask=mask)
    print(f"{cfg.aggregation.value:>24}  loss={out.loss.item():+.6f}")

Ten algorithms ship as registry entries. Every constant is traced to its paper or the paper's released code, with explicit notes where the paper is silent:

from polgrad import ALGORITHMS, describe_algorithm
print(describe_algorithm("dapo"))     # clip-higher eps, token-mean, no KL, with sources
print(sorted(ALGORITHMS))
# ['cispo', 'dapo', 'dr_grpo', 'grpo', 'grpo_tis', 'gspo', 'gspo_token', 'ppo', 'reinforce_pp', 'rloo']

polgrad.aggregate.effective_token_weights returns the closed-form per-token gradient weight each aggregation mode induces. That includes the micro-batch and gradient-accumulation weights, so their inequivalence to full-batch aggregation is computed exactly instead of being discovered in production.

Diagnostics

Metrics for the failure modes RL post-training actually exhibits, computed from tensors a training loop already has. Every threshold has a documented null distribution and a Monte Carlo test that calibrates it:

from polgrad.diagnostics import importance_ess, logprob_mismatch, clip_report, entropy_trend

ess = importance_ess(trainer_logprobs, rollout_logprobs, mask)   # ESS/n, with exact null == 1
gap = logprob_mismatch(trainer_logprobs, rollout_logprobs, mask) # k1/k2/k3 drift, catastrophic tokens
print(ess.summary()); print(gap.summary())

Entropy diagnostics accept either sampled-token logprobs (the Monte Carlo estimator, valid on-policy) or exact per-token entropies computed from full logits. The docs derive when each is valid. examples/diagnose_run.py synthesizes a 200-step run with injected pathologies (rollout↔trainer drift, entropy collapse, length bias) and shows each detector firing at the injection point.

Verification as public API

polgrad.verify exports the harness the library is tested with, so a new variant gets falsifiable checks for free: fp64 gradcheck over ragged batches, central finite-difference comparison against a supplied analytic formula, which fails on a wrong derivation as well as on wrong code, a softmax bandit with closed-form policy gradient and KL, and hand-derived golden cases whose arithmetic is written out in docs/derivations/goldens.md.

Conformance

src/polgrad/conformance/_vendor/ contains loss functions vendored verbatim from verl and OpenRLHF at pinned commits, with SHA256-enforced provenance headers; TRL is represented by a labeled reimplementation pinned to a version. Recorded fixtures keep CI framework-free. Demonstrated differences are registered in polgrad.conformance.DEVIATIONS, and each entry carries the pytest node id that demonstrates it. The wording is neutral by policy: an entry describes a difference between a published equation and a shipped default. Entries are not bug reports. Frameworks make defensible engineering choices; the problem is that the choices are currently invisible to the person training a model. docs/conformance.md documents the mechanism: what a comparison certifies, every registered target and deviation, the testing API, and drift detection.

from polgrad.conformance import DEVIATIONS, deviation_report
for d in DEVIATIONS:
    print(f"{d.framework} {d.version}: {d.description}")

A weekly scheduled job re-fetches the tracked upstream loss functions at HEAD, diffs them against the pins (tools/check_upstream_drift.py), and files an issue when a tracked definition has changed. Findings land in CHANGELOG.md under Upstream findings. The first one is already recorded: TRL's GRPOTrainer._compute_loss changed after the pinned v1.8.0, and the three loss-type branches polgrad models are byte-identical, so conformance results are unaffected.

For framework and trainer authors, polgrad.testing turns conformance into a test:

from polgrad.testing import assert_conforms

def test_my_grpo_matches_the_paper():
    assert_conforms(my_loss_fn, "grpo")   # seeded batches; compares loss and gradients

It installs as a pytest plugin, so a polgrad_batches fixture ships with the package. When semantics drift, assert_conforms raises with the full deviation report: maximum loss and gradient differences, and the worst-case seed.

Every claim is machine-checked

The derivation pages in docs/ state each identity and link the test that enforces it. A selection:

claim derivation enforced by
aggregate(x, m, mode)Σ wₜ xₜ with closed-form weights, bitwise, incl. autograd aggregation.md tests/test_aggregate.py
micro-batch aggregation weights == autograd of an explicit accumulation loop aggregation.md tests/test_cross.py
E[k1] = E[k3] = KL(π‖ref) exactly; k2's bias quantified kl.md tests/test_kl.py
k2-as-loss gradient ≡ unbiased reverse-KL score-function gradient (bitwise, all aggregations) kl.md tests/test_cross.py
k3-as-loss gradient ≠ ∇KL (Monte Carlo, >10 CLT tolerances from the analytic gradient) kl.md tests/test_kl.py
PG_CLIP ≡ PG ≡ REINFORCE gradients on-policy (bitwise) losses.md tests/test_losses.py
GSPO-token value ≡ GSPO-sequence value; gradient is token-local losses.md tests/test_losses.py
clip zero-gradient region == autograd-zero tokens, incl. dual-clip clipping.md tests/test_cross.py
RLOO leave-one-out ≡ (G/(G−1))·(r − mean) advantages.md tests/test_advantages.py
ESS/n null: ≡ 1 on-policy; → exp(−σ²) under N(0,σ²) log-weight drift ess.md tests/test_diagnostics_ess.py
entropy changepoint test has exact permutation level ≤ α entropy.md tests/test_diagnostics_entropy.py
all ten registry algorithms optimize a bandit to convergence variants.md tests/test_cross.py

The suite is 554 tests: analytic goldens, Hypothesis property tests over ragged masked batches, fp64 gradcheck, and seeded Monte Carlo with CLT-derived tolerances. It runs in about fourteen seconds serially on a laptop CPU (measured 13.8s; a few seconds with pytest -n auto, the CI invocation).

Scope

polgrad is not a trainer, a serving stack, a benchmark, or a critique of any framework. It never runs your training loop and never loads a model. It computes losses, gradients, advantages, KL estimates, and diagnostics from tensors you pass in.

References

Papers whose semantics are implemented and tested: PPO (1707.06347), GAE (1506.02438), dual-clip PPO (1912.09729), GRPO / DeepSeekMath (2402.03300), RLOO (2402.14740), REINFORCE++ (2501.03262), Dr.GRPO (2503.20783), DAPO (2503.14476), CISPO / MiniMax-M1 (2506.13585), GSPO (2507.18071); KL estimators after Schulman's approximation note, with the as-loss gradient analysis of 2512.21852 and 2510.01555.

License

Apache-2.0. Vendored framework code retains its upstream Apache-2.0 attribution; see NOTICE.

Project details


Download files

Download the file for your platform. If you're not sure which to choose, learn more about installing packages.

Source Distribution

polgrad-0.2.1.tar.gz (280.8 kB view details)

Uploaded Source

Built Distribution

If you're not sure about the file name format, learn more about wheel file names.

polgrad-0.2.1-py3-none-any.whl (104.0 kB view details)

Uploaded Python 3

File details

Details for the file polgrad-0.2.1.tar.gz.

File metadata

  • Download URL: polgrad-0.2.1.tar.gz
  • Upload date:
  • Size: 280.8 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? No
  • Uploaded via: uv/0.11.18 {"installer":{"name":"uv","version":"0.11.18","subcommand":["publish"]},"python":null,"implementation":{"name":null,"version":null},"distro":{"name":"macOS","version":null,"id":null,"libc":null},"system":{"name":null,"release":null},"cpu":null,"openssl_version":null,"setuptools_version":null,"rustc_version":null,"ci":null}

File hashes

Hashes for polgrad-0.2.1.tar.gz
Algorithm Hash digest
SHA256 793953946342b58075d5894121ead02d23f73a041277ebb672944c7acb32a8f7
MD5 d9d2228ab720de91f0009e50aec0cbcf
BLAKE2b-256 204f439b7315fb2cf0e5ef56ef02518365468e58e167c6363f9ba3f6e8f90f79

See more details on using hashes here.

File details

Details for the file polgrad-0.2.1-py3-none-any.whl.

File metadata

  • Download URL: polgrad-0.2.1-py3-none-any.whl
  • Upload date:
  • Size: 104.0 kB
  • Tags: Python 3
  • Uploaded using Trusted Publishing? No
  • Uploaded via: uv/0.11.18 {"installer":{"name":"uv","version":"0.11.18","subcommand":["publish"]},"python":null,"implementation":{"name":null,"version":null},"distro":{"name":"macOS","version":null,"id":null,"libc":null},"system":{"name":null,"release":null},"cpu":null,"openssl_version":null,"setuptools_version":null,"rustc_version":null,"ci":null}

File hashes

Hashes for polgrad-0.2.1-py3-none-any.whl
Algorithm Hash digest
SHA256 82dc05b973971b0dc043ef3de16670dfbcce1e6a7995ace5731a70e2113d7765
MD5 005f22b86be95987d1d102e950f510bc
BLAKE2b-256 20f2650dc261a131c52b67399748c49e308659d237aa476fcf8f803f5ac0ecd2

See more details on using hashes here.

Supported by

AWS Cloud computing and Security Sponsor Datadog Monitoring Depot Continuous Integration Fastly CDN Google Download Analytics Pingdom Monitoring Sentry Error logging StatusPage Status page