#!/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())