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