Skip to main content

multireward-grpo

Decoupled & conditioned multi-reward GRPO — advantage estimators, a generalized trainer, and the Theorem-3 verification harness from the paper "Decoupling and Conditioning Reshape Influence Allocation and the Gradient-Noise Floor in Multi-Reward GRPO under a Finite-Sample U-Statistic Analysis."

This package modularizes the experiment code so you can train your own multi-reward GRPO models, verify the correlation-aware MSE law on your own rollouts, and (optionally) run the whole thing on a cloud GPU.

pip install multireward-grpo            # core (numpy/scipy): advantage + analysis
pip install "multireward-grpo[llm]"     # + torch/transformers/peft: training & real generation
pip install "multireward-grpo[viz]"     # + matplotlib: the money-plot figure
pip install "multireward-grpo[data]"    # + datasets/hf-hub: dataset loaders & model push
pip install "multireward-grpo[runpod]"  # + requests: cloud GPU orchestration

The two orderings

Given a group of m rollouts each scored on R reward channels with weights w:

  • AN — Aggregate-then-Normalize (classic GRPO baseline): scalarize s = wᵀr, then group-normalize. The high-variance channel dominates (Prop 1) and the advantage resolution collapses under heterogeneous scales (Prop 2).
  • NA — Normalize-then-Aggregate (the decoupled estimator, = MO-GRPO/GDPO): group-normalize each channel, then take the weighted sum. Restores weight-proportional influence and gives the correlation-aware gradient-MSE floor (τ²/m)·wᵀCw (Theorem 3).
import numpy as np
from multireward_grpo import compute_advantage

rewards = np.array([[1.0, 0.3, 1.0],   # (m=4 rollouts, R=3 channels)
                    [0.0, 0.9, 1.0],
                    [1.0, 0.1, 0.0],
                    [0.0, 0.5, 1.0]])
w = np.array([1.0, 1.0, 0.5])
A_na = compute_advantage(rewards, w, mode="na")   # recommended
A_an = compute_advantage(rewards, w, mode="an")   # GRPO baseline

Train your own model

Bring your own prompts and your own reward function; the trainer runs group-relative policy optimization with a KL anchor and saves a LoRA adapter.

from multireward_grpo import GRPOConfig, GRPOTrainer

# prompts: list of strings, chat-message lists, or dicts with metadata
prompts = ["Write a polite refusal to a refund demand.", ...]

# reward_fn(completion, prompt) -> R channel scores (len == len(weights))
def reward_fn(completion, prompt):
    return (compliance(completion), politeness(completion), action(completion))

cfg = GRPOConfig(model="Qwen/Qwen2.5-1.5B-Instruct",
                 mode="na", weights=(1.0, 1.0, 0.5), n_steps=200, m=8)
history = GRPOTrainer(cfg, reward_fn, prompts).train()

Data format

Input Shape / type Notes
prompts list[str | list[dict] | dict] a string (user msg), chat messages [{"role","content"}], or {"prompt": ..., "gold": ...} with metadata passed through to the reward fn
reward_fn(completion, prompt) returns Sequence[float] of length R one score per reward channel; channel 0 is the gate for conditioning
weights tuple[float, ...] length R objective weights w
mode "na" | "an" | "single" na is the paper's recommendation

Reward tensors for the analysis tools use shape (P, K, m, R) = prompts × seeds × rollouts × reward channels.

Ready-made examples

from multireward_grpo.examples import FintechRewardFunction, make_fintech_prompts
from multireward_grpo import GRPOConfig, GRPOTrainer

prompts = make_fintech_prompts(400, seed=0)
cfg = GRPOConfig(mode="na", weights=(1.0, 1.0, 0.5))
GRPOTrainer(cfg, FintechRewardFunction(), prompts).train()

multireward_grpo.examples.gsm8k provides GSM8K loaders paired with multireward_grpo.rewards.MathRewardFunction (correctness / length / format).

Verify Theorem 3 on your rollouts

from multireward_grpo import analyze, summary_print
from multireward_grpo.generation import MockBackend, run_corpus, pack_for_analysis
import numpy as np

C = np.array([[1, 0.5, 0], [0.5, 1, 0], [0, 0, 1]])   # reward correlation
corpus = run_corpus(MockBackend(C=C), [(f"p{i}", "0") for i in range(40)],
                    m_grid=[8], K_seeds=200)
rewards = pack_for_analysis(corpus, m=8)               # (P, K, m, R)
result = analyze(rewards, w=np.array([1.0, 1.0, 0.5]))
summary_print(result)

Or from the shell:

multireward-grpo thm3-check --rho 0.5      # CPU, no GPU
multireward-grpo train --mode na --n-steps 50   # needs [llm] + GPU

Run on a cloud GPU (RunPod)

from multireward_grpo.runpod import RunPodClient
client = RunPodClient()  # reads RUNPOD_API_KEY from env or .env
client.run_command('pip install "multireward-grpo[llm]" && multireward-grpo train --mode na',
                   wall_clock_cap=1800)

Released artifacts (Hugging Face)

Datasets and fine-tuned models from the paper live under the eagle0504 namespace:

Datasets

Models (LoRA adapters for Qwen2.5-1.5B-Instruct)

Citation

If you use this package, please cite the paper (see the GitHub repository for the current BibTeX entry).

License

MIT

Metadata

Release files for multireward-grpo 0.1.1

For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.

Source distribution (sdist)

Source distribution for multireward-grpo 0.1.1
File Size Uploaded
multireward_grpo-0.1.1.tar.gz 386.6 kB Details

Built distribution (wheel)

Table of built distributions (wheels) for multireward-grpo 0.1.1
File Interpreter ABI Platform
multireward_grpo-0.1.1-py3-none-any.whl Python 3 none any Details

Total release size: 420.5 kB

Release files / multireward_grpo-0.1.1.tar.gz

Download URL multireward_grpo-0.1.1.tar.gz
Size 386.6 kB
Tags Source
SHA-256 checksum
How to use checksums
f01e77f5f356e48591894516a3eb409a3a8c9cdd78ce87a86ae74543fc145a53
BLAKE2b-256 checksum
How to use checksums
29ca182e82bd2874958370b90919c6ae15f05ab526df004acbbf97fab64db8bb
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
No
Uploaded via uv/0.12.5 {"installer":{"name":"uv","version":"0.12.5","subcommand":["publish"]},"python":null,"implementation":{"name":null,"version":null},"distro":{"name":"Ubuntu","version":"22.04","id":"jammy","libc":null},"system":{"name":null,"release":null},"cpu":null,"openssl_version":null,"setuptools_version":null,"rustc_version":null,"ci":null}

Release files / multireward_grpo-0.1.1-py3-none-any.whl

Download URL multireward_grpo-0.1.1-py3-none-any.whl
Size 33.9 kB
Tags Python 3
SHA-256 checksum
How to use checksums
8913f538109d20359aa6fa15b66ece02364d67594fc38323a8ae961001838aa2
BLAKE2b-256 checksum
How to use checksums
4983219a62a4a98a68110234a80e8fd1b98059aa4dd6a8e7663ff254dcc032bb
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
No
Uploaded via uv/0.12.5 {"installer":{"name":"uv","version":"0.12.5","subcommand":["publish"]},"python":null,"implementation":{"name":null,"version":null},"distro":{"name":"Ubuntu","version":"22.04","id":"jammy","libc":null},"system":{"name":null,"release":null},"cpu":null,"openssl_version":null,"setuptools_version":null,"rustc_version":null,"ci":null}

Release history Release notifications | RSS feed

This release

0.1.1 This release

2 release files

0.1.0

2 release files

Anthropic, PBC Visionary sponsor Bloomberg Visionary sponsor Hudson River Trading Visionary sponsor Meta Visionary sponsor NVIDIA Visionary sponsor Microsoft Sustainability sponsor Depot Continuous Integration AWS Cloud computing and Security Sponsor Datadog Monitoring Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page