raya / training /common.py
cderinbogaz's picture
Add training kit: train your own System-1 model
48c8658 verified
Raw History Blame Contribute Delete
5.57 kB
"""Task and data loading shared by label.py, train.py, evaluate.py and export_onnx.py.
A *task* (``task.json``) fixes the label set and the questions the model learns to answer:
{
"labels": ["small_model", "medium_model", "frontier_model"],
"questions": [ <Laya question>, ... ], # 1+ phrasings of the same decision
"rubric": "..." # optional, used by label.py
}
Each question is a regular Laya question. A ``choice`` question's ``criteria`` keys must be
exactly the labels (any order: options are shuffled in training). A ``score`` question's
``criteria`` list is ordinal and must have one entry per label, in label order.
*Data* is JSON Lines, one example per line:
{"prompt": "hi!", "label": "small_model"} # one gold label
{"prompt": "...", "labels": ["medium_model", "frontier_model"]} # several annotators -> soft target
{"state": {"ticket": "...", "plan": "pro"}, "label": "..."} # any Laya state instead of a prompt
{"prompt": "...", "label": "...", "split": "val"} # optional fixed split
Rows labelled ``"exclude"`` (by any annotator) are skipped.
"""
from __future__ import annotations
import json
from pathlib import Path
from typing import Any
EXCLUDE = "exclude"
class TaskError(ValueError):
"""Raised when task.json or a data file is malformed."""
def load_task(path: str | Path) -> dict[str, Any]:
"""Read and validate a task file; returns it with ``questions`` checked against ``labels``."""
task = json.loads(Path(path).read_text(encoding="utf-8"))
labels = task.get("labels")
if not isinstance(labels, list) or len(labels) < 2 or len(set(labels)) != len(labels):
raise TaskError("'labels' must be a list of 2+ distinct strings")
if EXCLUDE in labels:
raise TaskError(f"'{EXCLUDE}' is reserved for skipping rows; rename that label")
questions = task.get("questions")
if not isinstance(questions, list) or not questions:
raise TaskError("'questions' must be a non-empty list of Laya questions")
for i, q in enumerate(questions):
where = f"questions[{i}]"
if q.get("type") == "choice":
crit = q.get("criteria")
keys = list(crit) if isinstance(crit, (dict, list)) else None
if keys is None or set(keys) != set(labels) or len(keys) != len(labels):
raise TaskError(f"{where}: a choice question's criteria keys must be exactly the labels")
elif q.get("type") == "score":
crit = q.get("criteria")
if not isinstance(crit, list) or len(crit) != len(labels):
raise TaskError(f"{where}: a score question needs one criterion per label, in label order")
else:
raise TaskError(f"{where}: type must be 'choice' or 'score'")
if not isinstance(q.get("instructions"), str) or not q["instructions"].strip():
raise TaskError(f"{where}: 'instructions' must be a non-empty string")
return task
def row_state(row: dict[str, Any]) -> Any:
"""The Laya state for a data row: its ``state`` if given, else ``{"prompt": prompt}``."""
if "state" in row:
return row["state"]
if "prompt" in row:
return {"prompt": row["prompt"]}
raise TaskError("each row needs a 'prompt' or a 'state'")
def row_votes(row: dict[str, Any]) -> list[str]:
"""All label votes on a row (``label`` and/or ``labels``)."""
votes = list(row.get("labels") or [])
if row.get("label") is not None:
votes.append(row["label"])
return votes
def load_rows(path: str | Path, labels: list[str], *, require_labels: bool = True) -> list[dict[str, Any]]:
"""Read a JSONL data file into rows with a soft ``target`` over ``labels``.
The target is each label's share of the votes, so two annotators who disagree give a 50/50
target. ``gold`` is the label when every vote agrees, else ``None`` (such rows still train,
but accuracy is only reported on rows with a gold label).
"""
rows = []
for n, line in enumerate(Path(path).read_text(encoding="utf-8").splitlines(), 1):
if not line.strip():
continue
try:
row = json.loads(line)
except json.JSONDecodeError as exc:
raise TaskError(f"{path}:{n}: not valid JSON ({exc.msg})") from exc
row_state(row) # validates prompt/state presence
row.setdefault("id", n)
votes = row_votes(row)
if EXCLUDE in votes:
continue
if not votes:
if require_labels:
raise TaskError(f"{path}:{n}: no 'label' or 'labels'")
rows.append(row)
continue
unknown = sorted(set(votes) - set(labels))
if unknown:
raise TaskError(f"{path}:{n}: unknown label(s) {unknown}; expected one of {labels}")
row["target"] = [votes.count(label) / len(votes) for label in labels]
row["gold"] = votes[0] if len(set(votes)) == 1 else None
rows.append(row)
if not rows:
raise TaskError(f"{path}: no usable rows")
return rows
def answer_probs(answer: dict[str, Any], question: dict[str, Any], labels: list[str]) -> list[float]:
"""Map a Laya answer's probabilities back to label order (score options are ordinal)."""
probs = answer.get("probabilities") or {}
if question["type"] == "choice":
return [float(probs.get(label, 0.0)) for label in labels]
return [float(probs.get(str(i), probs.get(i, 0.0))) for i in range(len(labels))]