xianglinyang's picture
Add files using upload-large-folder tool
274951a verified
Raw History Blame Contribute Delete
15.5 kB
"""Verifiable coding reward for the RL-drift study.
One reward definition, shared by every algorithm so the *objective* is identical
across PPO / GRPO / DPO (the algorithm is the only thing that varies):
- GRPO calls ``coding_reward`` directly (TRL's ``reward_funcs`` API).
- PPO in TRL is reward-*model* based, so its objective must come from an RM
trained on this same signal — see rl_training/README.md for the confound note.
- DPO never sees a reward at train time; its (chosen, rejected) pairs are built
from this reward via ``build_preference_pairs`` so it chases the same target.
The main dataset (nvidia/Nemotron-RL-coding-competitive_coding) is competitive
programming with **stdin/stdout** tests, so the primary path is ``run_io_tests``:
feed each input on stdin, compare normalized stdout to the expected output. The
older ``run_unit_tests`` (Python ``assert`` snippets) is kept for assert-style
sets (MBPP / AceCode). Reward = fraction of tests passed, in [0, 1].
SECURITY: this executes model-generated code. The subprocess + timeout here is a
speed bump, not a sandbox. Run training on a throwaway box or wrap execution in a
real sandbox (container / firejail / nsjail) before pointing it at any dataset.
"""
from __future__ import annotations
import os
import re
import subprocess
import sys
import tempfile
from concurrent.futures import ThreadPoolExecutor
from functools import partial
from pathlib import Path
_CODE_FENCE = re.compile(r"```(?:[a-zA-Z0-9_+-]*)\n(.*?)```", re.DOTALL)
# Test cases run as subprocesses, so threads parallelize them fine (the GIL is
# released while waiting). Too many workers makes CPU-bound solutions contend
# and can push borderline cases over their timeout; cap conservatively. This is
# the default pool size; the training path overrides it via reward_num_workers
# (make_coding_reward -> coding_reward's num_workers) so concurrency tracks the
# run config.
_MAX_WORKERS = int(os.environ.get("VERIFIER_MAX_WORKERS", str(min(32, os.cpu_count() or 8))))
# --- Shared reward definition (the study's single training objective) -----------
# r(x, y) = fraction of tests passed. Every arm's TRAINING signal must use the
# same reward, or the algorithm comparison is confounded:
# - GRPO trains on the online reward over TRAIN_REWARD_MAX_TESTS tests;
# - DPO trains on pairs built from that same reward (build_dpo_pairs.py defaults
# to these constants);
# - PPO trains a reward model on those same pairs.
# The only cheapening vs the full suite is the test COUNT; the timeout is the same
# as eval so a correct-but-slow solution is graded identically in both places.
# The capped subset is spread evenly across the suite (see _subsample_indices),
# NOT the first N, because suites are often ordered easy->hard and grading only the
# easy prefix is reward-hackable. The subset is deterministic per prompt, so every
# completion in a GRPO group is judged on the same tests.
# Measurement (DriftCadenceCallback, final scoring) uses the FULL suite for a
# higher-fidelity, post-hoc equal-reward comparison — that is reporting, not the
# training objective, so it legitimately differs from the training r.
TRAIN_REWARD_MAX_TESTS = 12
# Canonical per-test timeout, shared by every arm's reward and by eval so a
# correct-but-slow solution grades the same everywhere. Single source of truth:
# the configs (grpo/ppo.yaml) and scripts/build_dpo_data.sh mirror this value.
# 5s is comfortably above what a correct competitive solution needs.
REWARD_TIMEOUT = 5.0
def _subsample_indices(n: int, max_tests: int | None) -> list[int]:
"""Evenly-spaced test indices spanning [0, n-1] inclusive (endpoints kept).
Deterministic in ``n`` alone, so a prompt's reward is a stable function and
all completions to that prompt are graded on the same cases. Spreads across
the suite rather than taking a contiguous prefix, so an easy->hard ordering
can't be gamed by solving only the easy end.
"""
if max_tests is None or max_tests >= n:
return list(range(n))
if max_tests <= 1:
return [0]
return sorted({round(i * (n - 1) / (max_tests - 1)) for i in range(max_tests)})
def _map_checks(checks: list, max_workers: int | None = None) -> list[bool]:
"""Run zero-arg test-case callables, in parallel when there are several.
``max_workers`` defaults to ``_MAX_WORKERS``; the training reward threads the
run config's worker count through here so grading concurrency is tunable.
"""
workers = _MAX_WORKERS if max_workers is None else max_workers
if len(checks) <= 1 or workers <= 1:
return [check() for check in checks]
with ThreadPoolExecutor(max_workers=min(workers, len(checks))) as pool:
return list(pool.map(lambda check: check(), checks))
def extract_code(completion: str) -> str:
"""Pull the last fenced code block from a completion, else the raw text."""
blocks = _CODE_FENCE.findall(completion or "")
return blocks[-1].strip() if blocks else (completion or "").strip()
def _completion_text(completion) -> str:
"""TRL hands completions as str (plain) or [{'role','content'}] (chat)."""
if isinstance(completion, str):
return completion
return completion[-1]["content"]
def _run(code: str, stdin: str | None, timeout: float) -> tuple[int, str]:
"""Execute ``code`` as a script in a fresh process; return (returncode, stdout)."""
with tempfile.TemporaryDirectory(prefix="drift-verify-") as workdir:
script = Path(workdir) / "solution.py"
script.write_text(code)
try:
result = subprocess.run(
[sys.executable, str(script)],
input=stdin,
capture_output=True,
text=True,
timeout=timeout,
cwd=workdir,
)
return result.returncode, result.stdout
except subprocess.TimeoutExpired:
return -1, ""
def _normalize(text: str) -> str:
"""Canonicalize competitive-judge output: unify newlines, rstrip each line,
drop trailing blank lines. Avoids false negatives from CRLF / trailing space."""
text = text.replace("\r\n", "\n").replace("\r", "\n")
lines = [line.rstrip() for line in text.split("\n")]
while lines and lines[-1] == "":
lines.pop()
return "\n".join(lines)
def _io_case_passes(code: str, stdin: str, expected: str, timeout: float) -> bool:
rc, stdout = _run(code, stdin, timeout)
return rc == 0 and _normalize(stdout) == _normalize(expected)
def _unit_case_passes(code: str, test: str, timeout: float) -> bool:
rc, _ = _run(f"{code}\n\n{test}\n", None, timeout)
return rc == 0
def _case_checks(completion: str, verifier: dict, timeout: float, max_tests: int | None = None) -> list:
"""One zero-arg callable per test case of a ``coding_task_v1`` verifier.
``max_tests`` caps the cases via ``_subsample_indices`` (a deterministic,
spread-out subset) BEFORE building callables, so the training reward's cheap
budget flows through the same test-case-level parallel path as full grading.
"""
code = extract_code(completion) if completion else ""
verifier_type = verifier.get("type")
if verifier_type in {"io_tests", "reference_io_tests"}:
inputs, outputs = verifier["test_inputs"], verifier["test_outputs"]
if not code or not inputs:
return []
idx = _subsample_indices(len(inputs), max_tests)
return [partial(_io_case_passes, code, inputs[i], outputs[i], timeout) for i in idx]
if verifier_type == "unit_tests":
tests = verifier["tests"]
if not code or not tests:
return []
idx = _subsample_indices(len(tests), max_tests)
return [partial(_unit_case_passes, code, tests[i], timeout) for i in idx]
raise ValueError(f"Unsupported verifier type: {verifier_type!r}")
def run_io_tests(
code: str,
inputs: list[str],
outputs: list[str],
timeout: float = 10.0,
max_tests: int | None = None,
) -> float:
"""stdin/stdout verifier (Nemotron / CodeContests). Pass fraction over cases.
``max_tests`` caps how many cases are executed (``None`` = all). During
training this is set low (see make_coding_reward) so grading is cheap; eval
and DPO-pair building leave it ``None`` for full-fidelity pass rates.
"""
verifier = {"type": "io_tests", "test_inputs": inputs, "test_outputs": outputs}
return run_verifier(code, verifier, timeout, max_tests)
def run_unit_tests(
code: str,
tests: list[str],
timeout: float = 10.0,
max_tests: int | None = None,
) -> float:
"""assert-style verifier (MBPP / AceCode). Pass fraction over snippets."""
return run_verifier(code, {"type": "unit_tests", "tests": tests}, timeout, max_tests)
def run_verifier(code: str, verifier: dict, timeout: float = 10.0, max_tests: int | None = None) -> float:
"""Dispatch a completion against a canonical ``coding_task_v1`` verifier."""
checks = _case_checks(code, verifier, timeout, max_tests)
if not checks:
return 0.0
results = _map_checks(checks)
return sum(results) / len(checks)
def coding_reward(
completions,
verifier=None,
test_inputs=None,
test_outputs=None,
timeout: float = 10.0,
max_tests: int | None = None,
num_workers: int | None = None,
**kwargs,
) -> list[float]:
"""GRPO/RLOO-compatible reward function.
TRL passes ``completions`` plus every dataset column as keyword arguments;
canonical data uses a ``verifier`` column. Legacy top-level
``test_inputs``/``test_outputs`` are still accepted during migration.
``max_tests`` caps tests per completion (the shared training budget).
``num_workers`` sizes the grading pool; every (completion x test-case) pair
across the whole rollout batch is flattened into one pool so the batch grades
concurrently, not one completion at a time. Both are set by make_coding_reward
from the run config; the bare defaults preserve the original full-suite,
default-pool behavior.
"""
if verifier is None:
if test_inputs is None or test_outputs is None:
raise ValueError("coding_reward requires either verifier or test_inputs/test_outputs")
verifier = [
{"type": "io_tests", "test_inputs": ti, "test_outputs": to}
for ti, to in zip(test_inputs, test_outputs)
]
# Flatten every (completion, test case) pair into one worker pool so the
# whole rollout batch verifies concurrently, not one completion at a time.
# max_tests is applied per completion inside _case_checks (spread subset).
per_completion = [
_case_checks(_completion_text(c), v, timeout, max_tests) for c, v in zip(completions, verifier)
]
flat_results = iter(_map_checks([check for checks in per_completion for check in checks], num_workers))
rewards = []
for checks in per_completion:
results = [next(flat_results) for _ in checks]
rewards.append(sum(results) / len(checks) if checks else 0.0)
return rewards
def make_coding_reward(
timeout: float = REWARD_TIMEOUT,
max_tests: int | None = TRAIN_REWARD_MAX_TESTS,
num_workers: int | None = None,
) -> callable:
"""Bind training-time grading knobs onto ``coding_reward`` for TRL.
TRL calls the reward function with a fixed signature (no timeout / cap args),
so the training config's cheap-grading settings are injected here instead.
``num_workers`` of ``None`` or ``<=0`` resolves to ``os.cpu_count()``. The
returned callable keeps ``__name__ == 'coding_reward'`` because TRL uses it
to name the reward's logged metric column.
"""
if num_workers and num_workers > 0:
resolved_workers = num_workers
else:
# ``reward_num_workers: 0`` means "use the process default" in the
# training configs. Respect the run-level cap before falling back to
# the machine CPU count: a GRPO batch flattens completion x test-case
# checks, so an unconstrained os.cpu_count() can exhaust file
# descriptors while spawning verifier subprocesses.
env_workers = os.environ.get("VERIFIER_MAX_WORKERS")
try:
resolved_workers = int(env_workers) if env_workers else (os.cpu_count() or 1)
except ValueError as exc:
raise ValueError("VERIFIER_MAX_WORKERS must be an integer") from exc
if resolved_workers <= 0:
raise ValueError("VERIFIER_MAX_WORKERS must be positive")
def coding_reward_fn(completions, **kwargs):
return coding_reward(
completions,
timeout=timeout,
max_tests=max_tests,
num_workers=resolved_workers,
**kwargs,
)
coding_reward_fn.__name__ = "coding_reward"
return coding_reward_fn
def score_completions(
completions: list[str],
verifier: dict,
timeout: float = 10.0,
max_tests: int | None = None,
num_workers: int | None = None,
) -> list[tuple[float, str]]:
"""Score candidate solutions with a canonical verifier.
``max_tests`` must match the training reward's budget when building DPO pairs
/ PPO reward-model data, so every arm shares one objective (build_dpo_pairs.py
passes it). Left ``None`` (full suite) for eval/measurement.
Every (completion x test-case) pair is flattened into one worker pool so the
whole candidate set grades concurrently, not one completion at a time — this
fully uses VERIFIER_MAX_WORKERS during DPO pair generation. ``num_workers``
of ``None`` resolves to ``_MAX_WORKERS`` (the env-configured default).
"""
texts = [_completion_text(c) for c in completions]
per_completion = [_case_checks(t, verifier, timeout, max_tests) for t in texts]
flat_results = iter(_map_checks([check for checks in per_completion for check in checks], num_workers))
scored = []
for text, checks in zip(texts, per_completion):
results = [next(flat_results) for _ in checks]
scored.append((sum(results) / len(checks) if checks else 0.0, text))
return scored
def build_preference_pairs(
prompt: str,
completions: list[str],
verifier: dict,
timeout: float = 10.0,
margin: float = 0.5,
max_tests: int | None = None,
) -> dict | None:
"""Turn scored candidates for one prompt into a DPO ``(chosen, rejected)`` row.
Scores every candidate with the same I/O verifier, pairs best vs worst, and
keeps the pair only when ``best - worst >= margin`` (default 0.5, per the
study design) so ties/near-ties don't inject label noise into DPO.
"""
scored = score_completions(completions, verifier, timeout, max_tests=max_tests)
return build_preference_pair_from_scored(prompt, scored, margin=margin)
def build_preference_pair_from_scored(
prompt: str,
scored_completions: list[tuple[float, str]],
margin: float = 0.5,
) -> dict | None:
"""Build a DPO pair from already-scored ``(score, completion)`` candidates."""
scored = sorted(scored_completions)
if not scored:
return None
worst_score, worst = scored[0]
best_score, best = scored[-1]
if best_score - worst_score < margin:
return None
return {
"prompt": prompt,
"chosen": _completion_text(best),
"rejected": _completion_text(worst),
"chosen_reward": best_score,
"rejected_reward": worst_score,
}