rl-training-debug-artifacts / code /evaluate_heldout_reward.py
xianglinyang's picture
Add files using upload-large-folder tool
274951a verified
Raw History Blame Contribute Delete
14.3 kB
#!/usr/bin/env python3
"""Evaluate M0 and aligned RL adapters on one frozen coding test split.
The endpoint must expose the base model and any requested LoRA aliases through
the OpenAI-compatible completions API. Prompts are rendered through the same
helper used by training and checkpoint reward evaluation. Output is append-only
and resumable at the (model, record_id) level.
"""
from __future__ import annotations
import argparse
import hashlib
import json
import math
import random
import time
from collections import defaultdict
from pathlib import Path
from statistics import mean
from typing import Any
import requests
from transformers import AutoTokenizer
from rl_training.data.schema import read_coding_jsonl
from rl_training.prompts import render_coding_prompt
from rl_training.rewards import REWARD_TIMEOUT, coding_reward
def sha256(path: Path) -> str:
return hashlib.sha256(path.read_bytes()).hexdigest()
def read_jsonl(path: Path) -> list[dict[str, Any]]:
if not path.is_file():
return []
return [json.loads(line) for line in path.read_text().splitlines() if line.strip()]
def request_completions(
*,
base_url: str,
api_key: str | None,
model: str,
prompts: list[str],
seed: int,
max_tokens: int,
temperature: float,
timeout: int,
) -> list[dict[str, Any]]:
payload = {
"model": model,
"prompt": prompts,
"n": 1,
"max_tokens": max_tokens,
"temperature": temperature,
"seed": seed,
}
error: Exception | None = None
for attempt in range(3):
try:
headers = {"Authorization": f"Bearer {api_key}"} if api_key else {}
response = requests.post(
f"{base_url.rstrip('/')}/completions",
headers=headers,
json=payload,
timeout=timeout,
)
response.raise_for_status()
choices = response.json().get("choices") or []
ordered = sorted(choices, key=lambda item: int(item["index"]))
if len(ordered) != len(prompts):
raise RuntimeError(
f"Expected {len(prompts)} completions from {model}, got {len(ordered)}"
)
return [
{
"text": str(item.get("text") or ""),
"finish_reason": item.get("finish_reason"),
}
for item in ordered
]
except (requests.RequestException, RuntimeError, ValueError) as exc:
error = exc
if attempt < 2:
time.sleep(5 * (attempt + 1))
raise RuntimeError(f"Completion request failed for {model}: {error}")
def bootstrap_delta(
baseline: list[float], treatment: list[float], *, seed: int, draws: int
) -> dict[str, float]:
if len(baseline) != len(treatment) or not baseline:
raise ValueError("Paired bootstrap requires equal non-empty vectors")
differences = [right - left for left, right in zip(baseline, treatment)]
rng = random.Random(seed)
estimates = sorted(
mean(differences[rng.randrange(len(differences))] for _ in differences)
for _ in range(draws)
)
low_index = max(0, math.floor(0.025 * draws))
high_index = min(draws - 1, math.ceil(0.975 * draws) - 1)
return {
"mean_delta": mean(differences),
"bootstrap_95_low": estimates[low_index],
"bootstrap_95_high": estimates[high_index],
"bootstrap_draws": draws,
}
def analyze(rows: list[dict[str, Any]], model_names: list[str], draws: int) -> dict[str, Any]:
by_model: dict[str, dict[str, dict[str, Any]]] = defaultdict(dict)
for row in rows:
key = str(row["model"])
record_id = str(row["record_id"])
if record_id in by_model[key]:
raise ValueError(f"Duplicate result for {key}/{record_id}")
by_model[key][record_id] = row
baseline_name = model_names[0]
baseline_ids = set(by_model[baseline_name])
if not baseline_ids:
raise ValueError(f"No baseline results for {baseline_name}")
if any(set(by_model[name]) != baseline_ids for name in model_names):
raise ValueError("All models must cover exactly the same record IDs")
ordered_ids = sorted(baseline_ids)
metrics: dict[str, Any] = {}
for name in model_names:
rewards = [float(by_model[name][record_id]["reward"]) for record_id in ordered_ids]
metrics[name] = {
"n_prompts": len(rewards),
"mean_test_case_pass_fraction": mean(rewards),
"strict_full_pass_rate": mean(value == 1.0 for value in rewards),
"nonzero_reward_rate": mean(value > 0.0 for value in rewards),
}
model_rows = [by_model[name][record_id] for record_id in ordered_ids]
if all("probably_truncated" in row for row in model_rows):
truncated = [bool(row["probably_truncated"]) for row in model_rows]
finish_reasons: dict[str, int] = defaultdict(int)
for row in model_rows:
finish_reasons[str(row.get("finish_reason"))] += 1
nontruncated_rewards = [
reward for reward, is_truncated in zip(rewards, truncated) if not is_truncated
]
metrics[name]["truncation"] = {
"count": sum(truncated),
"rate": mean(truncated),
"finish_reasons": dict(sorted(finish_reasons.items())),
"nontruncated_n": len(nontruncated_rewards),
"nontruncated_mean_test_case_pass_fraction": (
mean(nontruncated_rewards) if nontruncated_rewards else None
),
"nontruncated_strict_full_pass_rate": (
mean(value == 1.0 for value in nontruncated_rewards)
if nontruncated_rewards
else None
),
}
if name != baseline_name:
baseline = [
float(by_model[baseline_name][record_id]["reward"])
for record_id in ordered_ids
]
metrics[name]["paired_vs_m0"] = bootstrap_delta(
baseline, rewards, seed=20260817, draws=draws
)
full_baseline = [
float(by_model[baseline_name][record_id]["reward"] == 1.0)
for record_id in ordered_ids
]
full_treatment = [float(value == 1.0) for value in rewards]
metrics[name]["paired_full_pass_vs_m0"] = bootstrap_delta(
full_baseline, full_treatment, seed=20260818, draws=draws
)
return {
"schema": "rl_heldout_reward_analysis_v1",
"baseline_model": baseline_name,
"models": model_names,
"metrics": metrics,
}
def parse_models(values: list[str]) -> list[tuple[str, str]]:
result = []
for value in values:
if "=" not in value:
raise ValueError(f"--model must be LABEL=SERVED_MODEL, got {value!r}")
label, served = value.split("=", 1)
if not label or not served:
raise ValueError(f"Invalid --model value: {value!r}")
result.append((label, served))
if len({label for label, _ in result}) != len(result):
raise ValueError("Model labels must be unique")
return result
def main() -> int:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--dataset", type=Path, required=True)
parser.add_argument("--base-model", type=Path, required=True)
parser.add_argument("--base-url", default="http://127.0.0.1:8110/v1")
parser.add_argument(
"--api-key-file",
type=Path,
default=None,
help="Optional bearer-token file; omit for an unauthenticated localhost endpoint",
)
parser.add_argument("--model", action="append", required=True)
parser.add_argument("--output", type=Path, required=True)
parser.add_argument("--summary", type=Path, required=True)
parser.add_argument("--batch-size", type=int, default=16)
parser.add_argument("--max-tokens", type=int, default=512)
parser.add_argument(
"--truncation-buffer-tokens",
type=int,
default=24,
help="Also flag completions this close to --max-tokens as probably truncated.",
)
parser.add_argument("--temperature", type=float, default=1.0)
parser.add_argument("--seed", type=int, default=42)
parser.add_argument("--request-timeout", type=int, default=1800)
parser.add_argument("--verifier-timeout", type=float, default=REWARD_TIMEOUT)
parser.add_argument("--verifier-workers", type=int, default=32)
parser.add_argument("--bootstrap-draws", type=int, default=10000)
parser.add_argument(
"--limit", type=int, default=None, help="Optional leading-record limit for smoke tests."
)
args = parser.parse_args()
models = parse_models(args.model)
if args.batch_size < 1 or args.bootstrap_draws < 1:
parser.error("--batch-size and --bootstrap-draws must be positive")
if args.truncation_buffer_tokens < 0 or args.truncation_buffer_tokens >= args.max_tokens:
parser.error("--truncation-buffer-tokens must be in [0, --max-tokens)")
records = read_coding_jsonl(args.dataset)
if args.limit is not None:
if args.limit < 1:
parser.error("--limit must be positive")
records = records[: args.limit]
tokenizer = AutoTokenizer.from_pretrained(
args.base_model, trust_remote_code=True, local_files_only=True
)
prompts = [render_coding_prompt(tokenizer, row["prompt"]) for row in records]
prompt_token_counts = [
len(tokenizer(prompt, add_special_tokens=False)["input_ids"]) for prompt in prompts
]
args.output.parent.mkdir(parents=True, exist_ok=True)
args.summary.parent.mkdir(parents=True, exist_ok=True)
existing = read_jsonl(args.output)
completed = {(row["model"], row["record_id"]) for row in existing}
api_key = args.api_key_file.read_text().strip() if args.api_key_file else None
if args.api_key_file and not api_key:
raise ValueError("API key file is empty")
with args.output.open("a") as handle:
for model_index, (label, served_model) in enumerate(models):
pending = [
index
for index, row in enumerate(records)
if (label, row["record_id"]) not in completed
]
for offset in range(0, len(pending), args.batch_size):
indices = pending[offset : offset + args.batch_size]
batch_prompts = [prompts[index] for index in indices]
request_seed = args.seed + indices[0]
completion_results = request_completions(
base_url=args.base_url,
api_key=api_key,
model=served_model,
prompts=batch_prompts,
seed=request_seed,
max_tokens=args.max_tokens,
temperature=args.temperature,
timeout=args.request_timeout,
)
verifiers = [records[index]["verifier"] for index in indices]
completions = [result["text"] for result in completion_results]
rewards = coding_reward(
completions,
verifier=verifiers,
timeout=args.verifier_timeout,
max_tests=None,
num_workers=args.verifier_workers,
)
for index, result, reward in zip(indices, completion_results, rewards):
completion = result["text"]
completion_token_count = len(
tokenizer(completion, add_special_tokens=False)["input_ids"]
)
finish_reason = result["finish_reason"]
probably_truncated = (
finish_reason == "length"
or completion_token_count
>= args.max_tokens - args.truncation_buffer_tokens
)
row = {
"schema": "rl_heldout_reward_completion_v1",
"dataset_sha256": sha256(args.dataset),
"model": label,
"served_model": served_model,
"record_id": records[index]["record_id"],
"dataset_index": index,
"request_seed": request_seed,
"temperature": args.temperature,
"max_tokens": args.max_tokens,
"prompt_token_count": prompt_token_counts[index],
"completion_token_count": completion_token_count,
"finish_reason": finish_reason,
"probably_truncated": probably_truncated,
"truncation_buffer_tokens": args.truncation_buffer_tokens,
"completion": completion,
"reward": float(reward),
}
handle.write(json.dumps(row, ensure_ascii=False) + "\n")
handle.flush()
print(
f"[{label}] {min(offset + len(indices), len(pending))}/{len(pending)}",
flush=True,
)
all_rows = read_jsonl(args.output)
labels = [label for label, _ in models]
summary = analyze(all_rows, labels, args.bootstrap_draws)
summary.update(
{
"dataset": str(args.dataset),
"dataset_sha256": sha256(args.dataset),
"dataset_rows": len(records),
"rollouts_per_prompt": 1,
"generation": {
"temperature": args.temperature,
"max_tokens": args.max_tokens,
"seed": args.seed,
},
"verifier": {
"timeout": args.verifier_timeout,
"max_tests": None,
},
}
)
args.summary.write_text(json.dumps(summary, indent=2) + "\n")
print(json.dumps(summary, indent=2))
return 0
if __name__ == "__main__":
raise SystemExit(main())