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()