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
Download code/evaluate_heldout_reward.py from HilaryTorn/rl-training-debug-artifacts: direct link, hf CLI and curl.
- Browser
- Download file 14.3 kB
-
https://huggingface.co/HilaryTorn/rl-training-debug-artifacts/resolve/main/code/evaluate_heldout_reward.py
- Command line
-
hf download hf://HilaryTorn/rl-training-debug-artifacts/code/evaluate_heldout_reward.py
-
curl -L -o evaluate_heldout_reward.py https://huggingface.co/HilaryTorn/rl-training-debug-artifacts/resolve/main/code/evaluate_heldout_reward.py
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()) | |