raya / training /label.py
cderinbogaz's picture
Add training kit: train your own System-1 model
48c8658 verified
Raw History Blame Contribute Delete
6.64 kB
"""Label your prompts with two (or more) independent LLM annotators.
Raya's labels were made this way: two different strong models (Claude Opus and Claude Sonnet)
read the same written rubric and labelled every prompt without seeing each other's answers.
Where they agree you get a clean label; where they disagree the example keeps both votes and
train.py learns a 50/50 target instead of a wrong hard label.
export OPENAI_API_KEY=... # or any OpenAI-compatible provider
python label.py --task task.example.json --data prompts.jsonl --out labelled.jsonl \\
--annotator <model-a> \\
--annotator <model-b>@https://api.anthropic.com/v1/#ANTHROPIC_API_KEY
An annotator is ``MODEL[@BASE_URL][#API_KEY_ENV]``: BASE_URL defaults to
https://api.openai.com/v1/ and API_KEY_ENV to OPENAI_API_KEY. Use models from different
families so their mistakes are independent. Input rows need a ``prompt`` (or ``state``); any
existing labels are ignored. Re-running resumes: rows already labelled in --out are skipped.
"""
from __future__ import annotations
import argparse
import json
import os
import re
import threading
import time
import urllib.error
import urllib.request
from concurrent.futures import ThreadPoolExecutor, as_completed
from pathlib import Path
from common import EXCLUDE, load_rows, load_task, row_state
DEFAULT_BASE = "https://api.openai.com/v1/"
def parse_annotator(spec: str) -> dict:
spec, _, key_env = spec.partition("#")
model, _, base = spec.partition("@")
return {"model": model, "base": (base or DEFAULT_BASE).rstrip("/") + "/",
"key_env": key_env or "OPENAI_API_KEY", "name": model}
def system_prompt(task: dict) -> str:
labels = task["labels"]
rubric = task.get("rubric") or "\n".join(
f"- {label}: {task['questions'][0]['criteria'][label]}" for label in labels
if isinstance(task["questions"][0].get("criteria"), dict))
return (
"You are an independent annotator building a training set.\n"
f"{task['questions'][0]['instructions']}\n\n{rubric}\n\n"
f'Use "{EXCLUDE}" only if the input is unusable (empty, gibberish, no discernible request).\n'
"Judge the input yourself; do not guess from keywords. "
f'Answer with JSON only: {{"label": one of {labels + [EXCLUDE]}}}'
)
def call(annotator: dict, system: str, user: str, labels: list[str], retries: int = 6) -> str:
key = os.environ.get(annotator["key_env"])
if not key:
raise SystemExit(f"set {annotator['key_env']} for annotator {annotator['model']}")
body = json.dumps({"model": annotator["model"], "temperature": 0, "max_tokens": 50,
"messages": [{"role": "system", "content": system}, {"role": "user", "content": user}]})
request = urllib.request.Request(annotator["base"] + "chat/completions", data=body.encode(), headers={
"Content-Type": "application/json", "Authorization": f"Bearer {key}", "x-api-key": key,
"anthropic-version": "2023-06-01"})
for attempt in range(retries):
try:
with urllib.request.urlopen(request, timeout=120) as response:
text = json.load(response)["choices"][0]["message"]["content"] or ""
match = re.search(r"\{.*?\}", text, re.S)
label = json.loads(match.group(0)).get("label") if match else None
if label in labels or label == EXCLUDE:
return label
raise ValueError(f"unexpected answer {text[:80]!r}")
except urllib.error.HTTPError as exc:
if exc.code not in (408, 409, 429) and exc.code < 500:
raise SystemExit(f"{annotator['model']}: HTTP {exc.code} {exc.read()[:200]!r}") from exc
error = exc
except (urllib.error.URLError, TimeoutError, ValueError, KeyError) as exc:
error = exc
time.sleep(min(60, 2 ** attempt))
raise RuntimeError(f"{annotator['model']} failed after {retries} attempts: {error}")
def main() -> None:
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
ap.add_argument("--task", required=True)
ap.add_argument("--data", required=True, help="JSONL of rows with a 'prompt' or 'state'")
ap.add_argument("--out", required=True, help="labelled JSONL (appended; re-run to resume)")
ap.add_argument("--annotator", action="append", required=True, help="MODEL[@BASE_URL][#API_KEY_ENV]; repeat")
ap.add_argument("--workers", type=int, default=8)
ap.add_argument("--max-chars", type=int, default=6000, help="truncate long inputs sent to annotators")
args = ap.parse_args()
task = load_task(args.task)
labels = task["labels"]
annotators = [parse_annotator(a) for a in args.annotator]
if len(annotators) < 2:
print("warning: one annotator gives hard labels only; two different models are recommended")
system = system_prompt(task)
rows = load_rows(args.data, labels, require_labels=False)
out = Path(args.out)
done = {json.loads(l)["id"] for l in out.read_text(encoding="utf-8").splitlines() if l.strip()} if out.exists() else set()
todo = [r for r in rows if r["id"] not in done]
print(f"{len(rows)} rows, {len(done)} already labelled, {len(todo)} to go, {len(annotators)} annotators")
lock = threading.Lock()
def label_row(row: dict) -> dict:
state = row_state(row)
user = state["prompt"] if isinstance(state, dict) and set(state) == {"prompt"} else json.dumps(state, ensure_ascii=False)
if len(user) > args.max_chars:
user = user[:args.max_chars] + " …[truncated]"
votes = {a["name"]: call(a, system, user, labels) for a in annotators}
clean = {k: v for k, v in row.items() if k not in ("label", "labels", "target", "gold")}
return dict(clean, labels=list(votes.values()), annotators=votes)
agree = labelled = 0
with ThreadPoolExecutor(args.workers) as pool, out.open("a", encoding="utf-8") as sink:
for future in as_completed(pool.submit(label_row, r) for r in todo):
result = future.result()
with lock:
sink.write(json.dumps(result, ensure_ascii=False) + "\n")
sink.flush()
labelled += 1
agree += len(set(result["labels"])) == 1
if labelled % 50 == 0:
print(f"{labelled}/{len(todo)} labelled, annotators agree on {agree / labelled:.0%}", flush=True)
if labelled:
print(f"done: {labelled} labelled, annotators agree on {agree / labelled:.0%}")
if __name__ == "__main__":
main()