Download training/label.py from TextCortex/raya: direct link, hf CLI and curl.
- Browser
- Download file 6.64 kB
-
https://huggingface.co/TextCortex/raya/resolve/main/training/label.py
- Command line
-
hf download hf://TextCortex/raya/training/label.py
-
curl -L -o label.py https://huggingface.co/TextCortex/raya/resolve/main/training/label.py
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() | |