File size: 6,642 Bytes
48c8658 | 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 | """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()
|