Instructions to use HilaryTorn/rl-training-debug-artifacts with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- PEFT
How to use HilaryTorn/rl-training-debug-artifacts with PEFT:
Task type is invalid.
- Notebooks
- Google Colab
- Kaggle
File size: 14,325 Bytes
274951a | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 | #!/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())
|