cderinbogaz commited on
Commit
48c8658
·
verified ·
1 Parent(s): 8654b9c

Add training kit: train your own System-1 model

Browse files

Label (two independent LLM annotators), train, evaluate and export (ONNX fp32 + block-wise int8) scripts generalised from Raya's training. Model files unchanged.

README.md CHANGED
@@ -52,6 +52,22 @@ Raya was trained on three routing questions: the minimal choice above, a detaile
52
  3-level difficulty score. Use one of those. Option order does not matter because options were shuffled in
53
  training. Raya serves through Laya's Jev-compatible HTTP server (`POST /v1/systemone`).
54
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
55
  ## CPU inference (ONNX)
56
 
57
  No GPU? `onnx/` has ONNX Runtime builds of Raya that run through Laya's own `ONNXAgent`, which gives the same answer format as `laya.Agent`. Use a **512-token** input budget, which is what Raya was trained on.
@@ -124,6 +140,8 @@ annotators, not from measured downstream answer quality.
124
  validation prompts, minimal choice), per-question temperature fitted on validation. The seed was also
125
  chosen on validation only.
126
  - **Compute:** one NVIDIA RTX A6000, ~6 minutes.
 
 
127
 
128
  ## Intended use and limitations
129
 
 
52
  3-level difficulty score. Use one of those. Option order does not matter because options were shuffled in
53
  training. Raya serves through Laya's Jev-compatible HTTP server (`POST /v1/systemone`).
54
 
55
+ ## Train your own
56
+
57
+ The code that trained Raya is in [`training/`](training/README.md), generalised so you can train a fast
58
+ decision model on your own data:
59
+
60
+ - **Any choice or score question:** LLM routing like Raya, ticket triage, intent detection, escalation.
61
+ - **Labelling included:** labels come from two independent LLM annotators, which is how Raya's labels
62
+ were made.
63
+ - **Adapt Raya:** start from Raya to fit it to your own traffic (`--base TextCortex/raya`).
64
+ - **Serving:** exports to ONNX for CPU serving, checked against PyTorch.
65
+
66
+ ```bash
67
+ pip install -r training/requirements.txt
68
+ python training/train.py --task training/task.example.json --data my_labelled.jsonl --out my-router
69
+ ```
70
+
71
  ## CPU inference (ONNX)
72
 
73
  No GPU? `onnx/` has ONNX Runtime builds of Raya that run through Laya's own `ONNXAgent`, which gives the same answer format as `laya.Agent`. Use a **512-token** input budget, which is what Raya was trained on.
 
140
  validation prompts, minimal choice), per-question temperature fitted on validation. The seed was also
141
  chosen on validation only.
142
  - **Compute:** one NVIDIA RTX A6000, ~6 minutes.
143
+ - **Code:** [`training/`](training/README.md) reproduces this procedure on your own data. The training
144
+ data itself is not published.
145
 
146
  ## Intended use and limitations
147
 
training/README.md ADDED
@@ -0,0 +1,215 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Train your own System-1 model
2
+
3
+ This folder is the recipe we used to train [Raya](../README.md), packaged so you can train your own
4
+ fast decision model on your own data in an afternoon.
5
+
6
+ A **System-1 model** answers one well-defined question about an input, instantly and with calibrated
7
+ probabilities: *which model should answer this prompt?*, *does this ticket need a human?*, *which team
8
+ owns this request?* It is a small encoder (a [Laya](https://huggingface.co/convaiinnovations/laya)
9
+ decision model), not an LLM. It runs in tens of milliseconds on a CPU, costs nothing per call, and you
10
+ can host it anywhere.
11
+
12
+ ## The whole pipeline
13
+
14
+ ```bash
15
+ pip install -r requirements.txt
16
+
17
+ # 1. Describe your task: the labels and a few ways of asking the question
18
+ cp task.example.json my_task.json
19
+
20
+ # 2. Label your inputs with two independent LLMs (skip if you already have labels)
21
+ python label.py --task my_task.json --data prompts.jsonl --out labelled.jsonl \
22
+ --annotator <model-a> --annotator <model-b>@https://api.anthropic.com/v1/#ANTHROPIC_API_KEY
23
+
24
+ # 3. Train
25
+ python train.py --task my_task.json --data labelled.jsonl --out my-model
26
+
27
+ # 4. Check it on data it has never seen
28
+ python evaluate.py --model my-model --task my_task.json --data test.jsonl
29
+
30
+ # 5. (optional) Export for fast CPU serving, then publish
31
+ pip install -r requirements-onnx.txt
32
+ python export_onnx.py --model my-model --task my_task.json --data test.jsonl
33
+ python train.py ... --push-to-hub your-name/my-model # or upload my-model/ yourself
34
+ ```
35
+
36
+ **Try it in five minutes** with the bundled toy data (44 hand-written routing prompts, just enough to
37
+ see every step run; it is far too small to train a useful model):
38
+
39
+ ```bash
40
+ python train.py --task task.example.json --data data/example.jsonl --out my-router --epochs 2
41
+ python evaluate.py --model my-router --task task.example.json --data data/example_test.jsonl
42
+ ```
43
+
44
+ Then use it like any Laya model:
45
+
46
+ ```python
47
+ import json, laya
48
+
49
+ model = laya.Agent("my-router") # a local dir or a Hub repo id
50
+ question = json.load(open("my-router/task.json"))["questions"][0]
51
+ print(model.system_one({"prompt": "Prove that √2 is irrational."}, {"route": question})["answers"]["route"])
52
+ # {'choice': 'frontier_model', 'probabilities': {...}, ...}
53
+ ```
54
+
55
+ ## 1. Describe the task (`task.json`)
56
+
57
+ ```json
58
+ {
59
+ "labels": ["small_model", "medium_model", "frontier_model"],
60
+ "questions": [
61
+ {"type": "choice", "instructions": "Route this prompt to a model.",
62
+ "criteria": {"small_model": "simple requests", "medium_model": "moderately complex requests",
63
+ "frontier_model": "very hard requests"}},
64
+ {"type": "score", "instructions": "How difficult is this prompt for an AI assistant to answer well?",
65
+ "criteria": ["simple: a small model answers it perfectly", "moderate: needs a capable general model",
66
+ "hard: needs the strongest frontier model"]}
67
+ ],
68
+ "rubric": "Detailed labelling instructions for label.py (optional)."
69
+ }
70
+ ```
71
+
72
+ - **`labels`**: the possible answers, 2 or more.
73
+ - **`questions`**: one or more *phrasings* of the same decision, written as normal Laya questions. The
74
+ model trains on all of them, so it learns the decision rather than one exact wording. Raya was
75
+ trained on three phrasings and scores 80–81% on each.
76
+ - A `choice` question's `criteria` keys must be exactly your labels. Their order doesn't matter,
77
+ because options are shuffled during training.
78
+ - A `score` question is ordinal: one criterion per label, **in label order** (lowest first).
79
+ - **`rubric`**: what `label.py` shows the annotators. Be specific, with examples per label: the model
80
+ can only be as consistent as its labels.
81
+
82
+ `task.example.json` is Raya's exact routing task.
83
+
84
+ ## 2. Get labels (`label.py`)
85
+
86
+ Your data is JSON Lines, one input per line:
87
+
88
+ ```json
89
+ {"prompt": "hi there!", "label": "small_model"}
90
+ {"prompt": "Review this contract …", "labels": ["medium_model", "frontier_model"]}
91
+ {"state": {"ticket": "You charged me twice!", "plan": "enterprise"}, "label": "human_agent"}
92
+ {"prompt": "…", "label": "…", "split": "val"}
93
+ ```
94
+
95
+ - `label` is one gold answer. `labels` holds several annotators' votes: disagreements become soft
96
+ targets (50/50 above), which is better than forcing a hard label on a genuinely ambiguous input.
97
+ - Use `state` instead of `prompt` for structured inputs (any JSON object).
98
+ - Mark rows `"split": "val"` to fix your validation set; otherwise 10% is held out at random.
99
+ - Rows labelled `"exclude"` are skipped.
100
+
101
+ **No labels yet?** `label.py` has two (or more) different LLMs label every input independently from
102
+ your rubric, which is how Raya's labels were made (Claude Opus and Claude Sonnet, blind to each other).
103
+ It speaks the OpenAI chat-completions API, so it works with OpenAI, Anthropic's OpenAI-compatible
104
+ endpoint, OpenRouter, vLLM, Ollama and others. An annotator is `MODEL[@BASE_URL][#API_KEY_ENV]`.
105
+ It resumes where it stopped, and prints how often the annotators agree: that agreement rate is
106
+ roughly the ceiling your model can reach against these labels (78% for Raya's test set).
107
+
108
+ **How much data?**
109
+
110
+ - 1,000–5,000 is a good target for a first model.
111
+ - Raya used about 10,000.
112
+ - Use **real inputs** from your product where you can, in the languages you serve.
113
+ - Don't balance the classes artificially: `train.py` already up-weights rare labels.
114
+ - Keep a separate test set you never train or validate on.
115
+
116
+ ## 3. Train (`train.py`)
117
+
118
+ ```bash
119
+ python train.py --task my_task.json --data labelled.jsonl --out my-model
120
+ ```
121
+
122
+ What it does, in the same way as Raya's training:
123
+
124
+ - **Soft targets:** each label's share of the annotator votes, learned with cross-entropy.
125
+ - **Every question phrasing:** trained on each one, with choice options shuffled every epoch.
126
+ - **Class weights:** about 1/√(label frequency), so rare labels still count.
127
+ - **Memory:** the token-embedding table stays frozen.
128
+ - **Best epoch:** picked on validation accuracy (or `--select nll`).
129
+ - **Calibration:** a temperature is then fitted per question on validation, so that 0.9 means about 90%.
130
+
131
+ The output folder is a normal Laya checkpoint, plus `task.json` and `training_log.json`.
132
+
133
+ **Pick a starting point:**
134
+
135
+ | Flag | Starts from | When |
136
+ |---|---|---|
137
+ | *(default)* | Laya multilingual (mmBERT-base, 300M) | Most tasks, multilingual input (Raya was tested on 14 languages) |
138
+ | `--base TextCortex/raya` | Raya | LLM routing on your own traffic: adapt Raya instead of starting over |
139
+ | `--subfolder .` | Laya English (ModernBERT-large) | English-only input; a larger encoder, so slower |
140
+ | `--encoder <hf-id>` | Any Hugging Face encoder, fresh decision head | e.g. a larger multilingual encoder |
141
+ | `--base <dir or repo>` | Any Laya checkpoint | Continue from a model you trained before |
142
+
143
+ A bigger encoder was the biggest single lever in our experiments: with the same data, a large encoder
144
+ reached about 84% on Raya's benchmark where mmBERT-base topped out around 81–82%. Adding more data
145
+ barely moved the smaller model.
146
+
147
+ **Hardware:**
148
+
149
+ - **GPU:** Raya trained in about 6 minutes on one 48 GB RTX A6000 (batch 32, bf16). With the default
150
+ batch size of 16 it should fit on a 24 GB GPU.
151
+ - **Apple Silicon or CPU:** fine for a few thousand examples. `train.py` picks CUDA, then MPS, then CPU
152
+ automatically, and turns on gradient checkpointing off-GPU to save memory.
153
+
154
+ Useful flags: `--epochs` (3), `--batch-size` (16), `--lr-encoder` (2e-5), `--lr-head` (1e-4, or 3e-4
155
+ for a fresh head), `--max-tokens` (512), `--val-data`, `--seed`. Raya used these learning rates and
156
+ token budget with `--batch-size 32`, and its seed was chosen on validation only.
157
+
158
+ ## 4. Evaluate (`evaluate.py`)
159
+
160
+ ```bash
161
+ python evaluate.py --model my-model --task my_task.json --data test.jsonl [--out predictions.jsonl]
162
+ ```
163
+
164
+ For each question phrasing this prints accuracy, macro-F1 and a confusion matrix on rows with a single
165
+ gold label. It also prints the "always answer the most common label" baseline, which is the number to
166
+ beat, and per-decision latency. Pass `--onnx <file>` to score an exported model.
167
+
168
+ ## 5. Export and serve (`export_onnx.py`)
169
+
170
+ ```bash
171
+ python export_onnx.py --model my-model --task my_task.json --data test.jsonl
172
+ ```
173
+
174
+ This writes `my-model/onnx/model.onnx` (fp32) and `model-int8-blockwise.onnx`, then checks both against
175
+ PyTorch on your data. The export fails if any choice changes or probabilities drift by more than 0.001
176
+ (fp32) or 0.05 (int8). Serve either with Laya's `ONNXAgent`, which uses the same call and answer format:
177
+
178
+ ```python
179
+ from laya.onnx_agent import ONNXAgent
180
+
181
+ model = ONNXAgent("my-model", onnx_path="my-model/onnx/model.onnx")
182
+ model.cfg["max_len"] = 512 # match --max-tokens
183
+ ```
184
+
185
+ **Which ONNX file?** `model.onnx` matches PyTorch on any CPU. The block-wise int8 file kept Raya's accuracy
186
+ and was about 10–15% faster on x86 CPUs with VNNI instructions (Intel Cascade Lake or Alder Lake and
187
+ newer, AMD Zen 4 and newer), but slower on CPUs without VNNI and on ARM. Measure on your own hardware
188
+ before choosing it. In our tests 8 CPU threads were faster than 16.
189
+
190
+ ## Tips
191
+
192
+ - **Evaluate on data you never trained on.** Always, and keep the split fixed when comparing runs.
193
+ - **Don't tune on the test set.** Choose epochs, seeds and ensembles on validation only.
194
+ - **Match the serving token budget to training** (`--max-tokens`, default 512).
195
+ - **Your labels are the ceiling.** If two good annotators agree only 75% of the time, a 90% score
196
+ means the model learned your annotator's quirks. Tighten the rubric first.
197
+ - **Calibrate your action threshold on real traffic.** Before acting automatically on a prediction
198
+ (for example, only escalate when p(frontier) > 0.6), set the threshold from what your real traffic
199
+ looks like.
200
+
201
+ ## Files
202
+
203
+ | File | Purpose |
204
+ |---|---|
205
+ | `task.example.json` | Raya's routing task, a template for yours |
206
+ | `common.py` | Task and data format, validation (read its docstring for the full format) |
207
+ | `label.py` | Label inputs with independent LLM annotators |
208
+ | `train.py` | Fine-tune and calibrate |
209
+ | `evaluate.py` | Accuracy, macro-F1, confusion, latency |
210
+ | `export_onnx.py` | ONNX export (fp32 + block-wise int8) with equivalence checks |
211
+ | `data/example.jsonl`, `data/example_test.jsonl` | Tiny hand-written demo data for training and evaluation (not Raya's training data) |
212
+ | `data/unlabelled.jsonl` | A few unlabelled prompts to try `label.py` on |
213
+
214
+ Tested with laya 0.3.20, torch 2.14, transformers 5.17 and onnxruntime 1.30. The training data behind
215
+ Raya is not published. This code is Apache-2.0, like Raya and Laya.
training/common.py ADDED
@@ -0,0 +1,125 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Task and data loading shared by label.py, train.py, evaluate.py and export_onnx.py.
2
+
3
+ A *task* (``task.json``) fixes the label set and the questions the model learns to answer:
4
+
5
+ {
6
+ "labels": ["small_model", "medium_model", "frontier_model"],
7
+ "questions": [ <Laya question>, ... ], # 1+ phrasings of the same decision
8
+ "rubric": "..." # optional, used by label.py
9
+ }
10
+
11
+ Each question is a regular Laya question. A ``choice`` question's ``criteria`` keys must be
12
+ exactly the labels (any order: options are shuffled in training). A ``score`` question's
13
+ ``criteria`` list is ordinal and must have one entry per label, in label order.
14
+
15
+ *Data* is JSON Lines, one example per line:
16
+
17
+ {"prompt": "hi!", "label": "small_model"} # one gold label
18
+ {"prompt": "...", "labels": ["medium_model", "frontier_model"]} # several annotators -> soft target
19
+ {"state": {"ticket": "...", "plan": "pro"}, "label": "..."} # any Laya state instead of a prompt
20
+ {"prompt": "...", "label": "...", "split": "val"} # optional fixed split
21
+
22
+ Rows labelled ``"exclude"`` (by any annotator) are skipped.
23
+ """
24
+
25
+ from __future__ import annotations
26
+
27
+ import json
28
+ from pathlib import Path
29
+ from typing import Any
30
+
31
+ EXCLUDE = "exclude"
32
+
33
+
34
+ class TaskError(ValueError):
35
+ """Raised when task.json or a data file is malformed."""
36
+
37
+
38
+ def load_task(path: str | Path) -> dict[str, Any]:
39
+ """Read and validate a task file; returns it with ``questions`` checked against ``labels``."""
40
+ task = json.loads(Path(path).read_text(encoding="utf-8"))
41
+ labels = task.get("labels")
42
+ if not isinstance(labels, list) or len(labels) < 2 or len(set(labels)) != len(labels):
43
+ raise TaskError("'labels' must be a list of 2+ distinct strings")
44
+ if EXCLUDE in labels:
45
+ raise TaskError(f"'{EXCLUDE}' is reserved for skipping rows; rename that label")
46
+ questions = task.get("questions")
47
+ if not isinstance(questions, list) or not questions:
48
+ raise TaskError("'questions' must be a non-empty list of Laya questions")
49
+ for i, q in enumerate(questions):
50
+ where = f"questions[{i}]"
51
+ if q.get("type") == "choice":
52
+ crit = q.get("criteria")
53
+ keys = list(crit) if isinstance(crit, (dict, list)) else None
54
+ if keys is None or set(keys) != set(labels) or len(keys) != len(labels):
55
+ raise TaskError(f"{where}: a choice question's criteria keys must be exactly the labels")
56
+ elif q.get("type") == "score":
57
+ crit = q.get("criteria")
58
+ if not isinstance(crit, list) or len(crit) != len(labels):
59
+ raise TaskError(f"{where}: a score question needs one criterion per label, in label order")
60
+ else:
61
+ raise TaskError(f"{where}: type must be 'choice' or 'score'")
62
+ if not isinstance(q.get("instructions"), str) or not q["instructions"].strip():
63
+ raise TaskError(f"{where}: 'instructions' must be a non-empty string")
64
+ return task
65
+
66
+
67
+ def row_state(row: dict[str, Any]) -> Any:
68
+ """The Laya state for a data row: its ``state`` if given, else ``{"prompt": prompt}``."""
69
+ if "state" in row:
70
+ return row["state"]
71
+ if "prompt" in row:
72
+ return {"prompt": row["prompt"]}
73
+ raise TaskError("each row needs a 'prompt' or a 'state'")
74
+
75
+
76
+ def row_votes(row: dict[str, Any]) -> list[str]:
77
+ """All label votes on a row (``label`` and/or ``labels``)."""
78
+ votes = list(row.get("labels") or [])
79
+ if row.get("label") is not None:
80
+ votes.append(row["label"])
81
+ return votes
82
+
83
+
84
+ def load_rows(path: str | Path, labels: list[str], *, require_labels: bool = True) -> list[dict[str, Any]]:
85
+ """Read a JSONL data file into rows with a soft ``target`` over ``labels``.
86
+
87
+ The target is each label's share of the votes, so two annotators who disagree give a 50/50
88
+ target. ``gold`` is the label when every vote agrees, else ``None`` (such rows still train,
89
+ but accuracy is only reported on rows with a gold label).
90
+ """
91
+ rows = []
92
+ for n, line in enumerate(Path(path).read_text(encoding="utf-8").splitlines(), 1):
93
+ if not line.strip():
94
+ continue
95
+ try:
96
+ row = json.loads(line)
97
+ except json.JSONDecodeError as exc:
98
+ raise TaskError(f"{path}:{n}: not valid JSON ({exc.msg})") from exc
99
+ row_state(row) # validates prompt/state presence
100
+ row.setdefault("id", n)
101
+ votes = row_votes(row)
102
+ if EXCLUDE in votes:
103
+ continue
104
+ if not votes:
105
+ if require_labels:
106
+ raise TaskError(f"{path}:{n}: no 'label' or 'labels'")
107
+ rows.append(row)
108
+ continue
109
+ unknown = sorted(set(votes) - set(labels))
110
+ if unknown:
111
+ raise TaskError(f"{path}:{n}: unknown label(s) {unknown}; expected one of {labels}")
112
+ row["target"] = [votes.count(label) / len(votes) for label in labels]
113
+ row["gold"] = votes[0] if len(set(votes)) == 1 else None
114
+ rows.append(row)
115
+ if not rows:
116
+ raise TaskError(f"{path}: no usable rows")
117
+ return rows
118
+
119
+
120
+ def answer_probs(answer: dict[str, Any], question: dict[str, Any], labels: list[str]) -> list[float]:
121
+ """Map a Laya answer's probabilities back to label order (score options are ordinal)."""
122
+ probs = answer.get("probabilities") or {}
123
+ if question["type"] == "choice":
124
+ return [float(probs.get(label, 0.0)) for label in labels]
125
+ return [float(probs.get(str(i), probs.get(i, 0.0))) for i in range(len(labels))]
training/data/example.jsonl ADDED
@@ -0,0 +1,44 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {"prompt": "Derive the Black-Scholes PDE from first principles using Itô's lemma and a delta-hedged portfolio.", "label": "frontier_model"}
2
+ {"prompt": "Write a two-line birthday wish for my sister.", "label": "small_model"}
3
+ {"prompt": "Write a cover letter for a senior data scientist role at a bank.", "labels": ["medium_model", "frontier_model"]}
4
+ {"prompt": "hi there!", "label": "small_model"}
5
+ {"prompt": "What are the pros and cons of remote work for a 50-person agency?", "label": "medium_model"}
6
+ {"prompt": "Rédige une lettre de motivation pour un poste de chef de projet marketing.", "label": "medium_model"}
7
+ {"prompt": "Ciao! Come stai?", "label": "small_model"}
8
+ {"prompt": "Give me a synonym for 'happy'.", "label": "small_model"}
9
+ {"prompt": "Write a Python function that parses a CSV of orders and returns total revenue per customer.", "label": "medium_model"}
10
+ {"prompt": "Create a one-week meal plan for a vegetarian who trains for a marathon.", "label": "medium_model"}
11
+ {"prompt": "Convert 5 miles to kilometers.", "label": "small_model"}
12
+ {"prompt": "How do I reverse a list in Python?", "label": "small_model"}
13
+ {"prompt": "thanks, that's all", "label": "small_model"}
14
+ {"prompt": "What's the capital of Canada?", "label": "small_model"}
15
+ {"prompt": "Draft a LinkedIn post announcing our new office in Munich.", "label": "medium_model"}
16
+ {"prompt": "Write a polite email to a client explaining that their order will arrive two weeks late.", "label": "medium_model"}
17
+ {"prompt": "¿Cuántos días tiene febrero en un año bisiesto?", "label": "small_model"}
18
+ {"prompt": "Explain quantum entanglement.", "labels": ["small_model", "medium_model"]}
19
+ {"prompt": "Refactor this React component to use hooks instead of a class: [80-line class component with state and lifecycle methods]", "label": "medium_model"}
20
+ {"prompt": "Build a discounted cash flow model for a SaaS company with 40% growth, negative free cash flow and a 2-year path to breakeven; justify every assumption.", "label": "frontier_model"}
21
+ {"prompt": "Scrivi una recensione di un ristorante giapponese immaginario, circa 300 parole.", "label": "medium_model"}
22
+ {"prompt": "Optimize this SQL query that takes 40 seconds on a 200M-row table: [query with three joins and a window function]", "labels": ["medium_model", "frontier_model"]}
23
+ {"prompt": "Find the race condition in this lock-free queue implementation and prove your fix is linearizable: [120 lines of Rust using atomics]", "label": "frontier_model"}
24
+ {"prompt": "Explain how HTTPS works to a non-technical colleague.", "label": "medium_model"}
25
+ {"prompt": "Écris un compilateur minimal pour un langage à la Lisp en Rust, avec analyse lexicale, analyse syntaxique et évaluation.", "label": "frontier_model"}
26
+ {"prompt": "Compare three treatment strategies for a patient with resistant hypertension and chronic kidney disease, citing the relevant trial evidence.", "label": "frontier_model"}
27
+ {"prompt": "Analysiere die kartellrechtlichen Risiken einer Fusion zweier Marktführer im deutschen Zementmarkt und schlage Abhilfemaßnahmen vor.", "label": "frontier_model"}
28
+ {"prompt": "Fix the typo: 'I recieved your mesage'", "label": "small_model"}
29
+ {"prompt": "Translate this paragraph into French: [one short paragraph]", "labels": ["small_model", "medium_model"]}
30
+ {"prompt": "What is 17 * 23?", "label": "small_model"}
31
+ {"prompt": "Bonjour, comment ça va ?", "label": "small_model"}
32
+ {"prompt": "Wie spät ist es in Tokio, wenn es in Berlin 9 Uhr ist?", "label": "small_model"}
33
+ {"prompt": "Develop a go-to-market strategy for entering the Japanese B2B software market, including channel partners, pricing and a 3-year plan.", "label": "frontier_model"}
34
+ {"prompt": "What does HTML stand for?", "label": "small_model"}
35
+ {"prompt": "Translate this product description into German and keep the marketing tone: [four-paragraph description of a smart thermostat]", "label": "medium_model"}
36
+ {"prompt": "Prove that there are infinitely many primes of the form 4k+3.", "label": "frontier_model"}
37
+ {"prompt": "Schreibe eine professionelle E-Mail an einen Kunden über eine Lieferverzögerung von zwei Wochen.", "label": "medium_model"}
38
+ {"prompt": "Escribe un artículo de blog de 600 palabras sobre consejos para ahorrar energía en casa.", "label": "medium_model"}
39
+ {"prompt": "Design a multi-region, active-active payment system that guarantees exactly-once settlement, and explain how you handle network partitions.", "label": "frontier_model"}
40
+ {"prompt": "Review this 40-page share purchase agreement and list every clause that shifts liability to the buyer, with the risk each one creates.", "label": "frontier_model"}
41
+ {"prompt": "Summarize the key points of this meeting transcript in bullet points: [transcript of a 30-minute product sync about launch timing, pricing and support staffing]", "label": "medium_model"}
42
+ {"prompt": "Translate 'good morning' into Spanish.", "label": "small_model"}
43
+ {"prompt": "Explain the difference between a stock and a bond with examples.", "label": "medium_model"}
44
+ {"prompt": "Write a SQL query that finds the top 5 products by revenue for each month of 2025.", "label": "medium_model"}
training/data/example_test.jsonl ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {"prompt": "hello!", "label": "small_model"}
2
+ {"prompt": "What is the boiling point of water in Fahrenheit?", "label": "small_model"}
3
+ {"prompt": "Translate 'thank you' into Japanese.", "label": "small_model"}
4
+ {"prompt": "Was ist 12 plus 30?", "label": "small_model"}
5
+ {"prompt": "Rewrite this sentence more formally: 'gonna be late, sorry'", "label": "small_model"}
6
+ {"prompt": "Write a short product update email for our customers about a new export feature.", "label": "medium_model"}
7
+ {"prompt": "Explain what a REST API is with a simple example.", "label": "medium_model"}
8
+ {"prompt": "Summarize the plot of Hamlet in one paragraph.", "label": "medium_model"}
9
+ {"prompt": "Écris un e-mail pour reporter une réunion à la semaine prochaine.", "label": "medium_model"}
10
+ {"prompt": "Write a JavaScript function that debounces another function.", "label": "medium_model"}
11
+ {"prompt": "Prove that the square root of 2 is irrational and generalize the argument to any non-square integer.", "label": "frontier_model"}
12
+ {"prompt": "Design the data model and consistency strategy for a global inventory system with offline-capable warehouses.", "label": "frontier_model"}
13
+ {"prompt": "Evaluate the tax implications of relocating a German GmbH's headquarters to the Netherlands, including exit taxation.", "label": "frontier_model"}
training/data/unlabelled.jsonl ADDED
@@ -0,0 +1,4 @@
 
 
 
 
 
1
+ {"prompt": "hello!"}
2
+ {"prompt": "What is the boiling point of water in Fahrenheit?"}
3
+ {"prompt": "Translate 'thank you' into Japanese."}
4
+ {"prompt": "Was ist 12 plus 30?"}
training/evaluate.py ADDED
@@ -0,0 +1,97 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Score a trained model on a labelled JSONL file: accuracy, macro-F1, confusion and latency.
2
+
3
+ python evaluate.py --model my-router --task task.example.json --data data/example.jsonl
4
+ python evaluate.py --model my-router --onnx my-router/onnx/model-int8-blockwise.onnx ...
5
+
6
+ Accuracy counts rows whose annotators all agreed (a single gold label). Always test on data the
7
+ model never trained or validated on. The majority-label baseline is printed for comparison.
8
+ """
9
+
10
+ from __future__ import annotations
11
+
12
+ import argparse
13
+ import json
14
+ import os
15
+ import time
16
+
17
+ os.environ.setdefault("USE_TF", "0")
18
+
19
+ import numpy as np
20
+
21
+ from common import answer_probs, load_rows, load_task, row_state
22
+
23
+
24
+ def load_model(path: str, onnx_path: str | None, device: str | None, max_tokens: int):
25
+ if onnx_path:
26
+ from laya.onnx_agent import ONNXAgent
27
+
28
+ model = ONNXAgent(path, onnx_path=onnx_path)
29
+ else:
30
+ import laya
31
+
32
+ model = laya.Agent(path, device=device)
33
+ model.cfg["max_len"] = max_tokens
34
+ return model
35
+
36
+
37
+ def main() -> None:
38
+ ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
39
+ ap.add_argument("--model", required=True, help="checkpoint dir or Hub repo id")
40
+ ap.add_argument("--task", required=True)
41
+ ap.add_argument("--data", required=True, help="labelled JSONL the model has not seen")
42
+ ap.add_argument("--onnx", help="score this ONNX file (with the checkpoint's config/tokenizer) instead")
43
+ ap.add_argument("--device", default=None)
44
+ ap.add_argument("--max-tokens", type=int, default=512)
45
+ ap.add_argument("--out", help="write per-row predictions to this JSONL")
46
+ args = ap.parse_args()
47
+
48
+ task = load_task(args.task)
49
+ labels, questions = task["labels"], task["questions"]
50
+ rows = [r for r in load_rows(args.data, labels) if r["gold"] is not None]
51
+ model = load_model(args.model, args.onnx, args.device, args.max_tokens)
52
+
53
+ preds = {qi: [] for qi in range(len(questions))}
54
+ latencies = []
55
+ out = open(args.out, "w", encoding="utf-8") if args.out else None
56
+ for row in rows:
57
+ record = {"id": row["id"], "gold": row["gold"], "predictions": {}}
58
+ for qi, q in enumerate(questions):
59
+ t = time.perf_counter()
60
+ answer = model.system_one(row_state(row), {"q": q})["answers"]["q"]
61
+ latencies.append((time.perf_counter() - t) * 1000)
62
+ probs = answer_probs(answer, q, labels)
63
+ preds[qi].append(int(np.argmax(probs)))
64
+ record["predictions"][qi] = dict(zip(labels, [round(p, 4) for p in probs]))
65
+ if out:
66
+ out.write(json.dumps(record, ensure_ascii=False) + "\n")
67
+ if out:
68
+ out.close()
69
+
70
+ golds = [labels.index(r["gold"]) for r in rows]
71
+ majority = max(labels, key=lambda l: sum(r["gold"] == l for r in rows))
72
+ print(f"{len(rows)} rows with a gold label; always answering {majority!r} scores "
73
+ f"{sum(r['gold'] == majority for r in rows) / len(rows):.1%}")
74
+ for qi, q in enumerate(questions):
75
+ p = preds[qi]
76
+ acc = np.mean([a == b for a, b in zip(p, golds)])
77
+ f1s = []
78
+ for k in range(len(labels)):
79
+ tp = sum(a == k == b for a, b in zip(p, golds))
80
+ fp = sum(a == k != b for a, b in zip(p, golds))
81
+ fn = sum(b == k != a for a, b in zip(p, golds))
82
+ f1s.append(2 * tp / (2 * tp + fp + fn) if tp else 0.0)
83
+ confusion = [[sum(g == i and a == j for a, g in zip(p, golds)) for j in range(len(labels))]
84
+ for i in range(len(labels))]
85
+ print(f"\nquestion {qi} ({q['type']}): {q['instructions'][:70]!r}")
86
+ print(f" accuracy {acc:.1%} macro-F1 {np.mean(f1s):.3f}")
87
+ print(" confusion (rows = gold, columns = predicted):")
88
+ width = max(len(l) for l in labels)
89
+ print(" " + " " * width + " " + " ".join(f"{l:>{width}}" for l in labels))
90
+ for label, line in zip(labels, confusion):
91
+ print(f" {label:>{width}} " + " ".join(f"{n:>{width}}" for n in line))
92
+ lat = np.array(latencies)
93
+ print(f"\nlatency per decision: p50 {np.percentile(lat, 50):.0f} ms, p95 {np.percentile(lat, 95):.0f} ms")
94
+
95
+
96
+ if __name__ == "__main__":
97
+ main()
training/export_onnx.py ADDED
@@ -0,0 +1,151 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Export a trained checkpoint to ONNX for fast CPU serving, and check it against PyTorch.
2
+
3
+ python export_onnx.py --model my-router --task task.example.json --data data/example.jsonl
4
+
5
+ Writes into ``<model>/onnx/``:
6
+
7
+ * ``model.onnx``: fp32, matches PyTorch on any CPU (the safe default);
8
+ * ``model-int8-blockwise.onnx``: 8-bit block-wise encoder weights (ONNX Runtime MatMulNBits,
9
+ block 32). It kept Raya's accuracy and ran ~10-15% faster on CPUs with VNNI int8 instructions
10
+ (Intel Cascade Lake/Alder Lake+, AMD Zen 4+), but slower without VNNI. Skip with --fp32-only.
11
+
12
+ Serve either with Laya's ONNXAgent (same answers and format as laya.Agent):
13
+
14
+ from laya.onnx_agent import ONNXAgent
15
+ agent = ONNXAgent("my-router", onnx_path="my-router/onnx/model.onnx"); agent.cfg["max_len"] = 512
16
+
17
+ Every exported file is checked on up to 20 rows of --data: the choice must not change and
18
+ probabilities must stay within 0.001 (fp32) or 0.05 (int8) of PyTorch.
19
+ """
20
+
21
+ from __future__ import annotations
22
+
23
+ import argparse
24
+ import json
25
+ import os
26
+ from pathlib import Path
27
+
28
+ os.environ.setdefault("USE_TF", "0")
29
+
30
+ import numpy as np
31
+ import onnx
32
+ import torch
33
+ from laya import Agent
34
+ from laya.onnx_agent import ONNXAgent
35
+
36
+ from common import answer_probs, load_rows, load_task, row_state
37
+
38
+ GRAPH_INPUTS = ("input_ids", "attention_mask", "marker_pos", "marker_mask", "qtype")
39
+
40
+
41
+ def capture_inputs(agent: Agent, state, question: dict) -> dict[str, np.ndarray]:
42
+ """The network inputs Laya builds for one decision (used as the export example)."""
43
+ captured: dict[str, np.ndarray] = {}
44
+ model = agent.model
45
+
46
+ class Capture(torch.nn.Module):
47
+ def forward(self, *tensors):
48
+ captured.update({n: t.numpy() for n, t in zip(GRAPH_INPUTS, tensors)})
49
+ return model(*tensors)
50
+
51
+ agent.model = Capture()
52
+ try:
53
+ agent.system_one(state, {"q": question})
54
+ finally:
55
+ agent.model = model
56
+ return captured
57
+
58
+
59
+ def export_fp32(agent: Agent, state, question: dict, path: Path) -> None:
60
+ example = capture_inputs(agent, state, question)
61
+ seq = torch.export.Dim("seq", min=8, max=4096)
62
+ opts = torch.export.Dim("opts", min=1, max=16)
63
+ torch.onnx.export(agent.model.eval(), tuple(torch.from_numpy(example[k]) for k in GRAPH_INPUTS), str(path),
64
+ input_names=list(GRAPH_INPUTS), output_names=["logits", "act_logits"],
65
+ dynamic_shapes=({1: seq}, {1: seq}, {1: opts}, {1: opts}, None),
66
+ dynamo=True, external_data=False)
67
+ # Shapes recorded by the exporter contradict ONNX shape inference during quantization;
68
+ # drop them (ONNX Runtime re-infers them at load).
69
+ model = onnx.load(str(path))
70
+ del model.graph.value_info[:]
71
+ onnx.save(model, str(path))
72
+
73
+
74
+ def encoder_weight_matmuls(path: Path, n_layers: int) -> list[str] | None:
75
+ """Names of the encoder's weight matmuls (4 per ModernBERT/mmBERT layer), or None if not found."""
76
+ model = onnx.load(str(path), load_external_data=False)
77
+ constants = {i.name for i in model.graph.initializer}
78
+ producers = {o: n for n in model.graph.node for o in n.output}
79
+
80
+ def is_constant(name: str) -> bool:
81
+ p = producers.get(name)
82
+ return name in constants or (p is not None and p.op_type in ("Transpose", "Cast", "Identity")
83
+ and all(is_constant(i) for i in p.input))
84
+
85
+ names = [node.name for node in model.graph.node
86
+ if node.op_type == "MatMul" and is_constant(node.input[1])
87
+ and "encoder.layers." in {p.key: p.value for p in node.metadata_props}.get("pkg.torch.onnx.name_scopes", "")]
88
+ return names if len(names) == 4 * n_layers else None
89
+
90
+
91
+ def quantize_blockwise(src: Path, dst: Path, nodes: list[str]) -> None:
92
+ from onnxruntime.quantization.matmul_nbits_quantizer import MatMulNBitsQuantizer
93
+
94
+ quantizer = MatMulNBitsQuantizer(onnx.load(str(src)), bits=8, block_size=32, is_symmetric=True,
95
+ accuracy_level=4, nodes_to_include=nodes)
96
+ quantizer.process()
97
+ quantizer.model.save_model_to_file(str(dst), use_external_data_format=False)
98
+
99
+
100
+ def check(agent: Agent, model_dir: str, onnx_path: Path, rows, question, labels, max_tokens, tolerance) -> float:
101
+ served = ONNXAgent(model_dir, onnx_path=str(onnx_path))
102
+ served.cfg["max_len"] = max_tokens
103
+ worst = 0.0
104
+ for row in rows:
105
+ state = row_state(row)
106
+ want = answer_probs(agent.system_one(state, {"q": question})["answers"]["q"], question, labels)
107
+ got = answer_probs(served.system_one(state, {"q": question})["answers"]["q"], question, labels)
108
+ delta = max(abs(a - b) for a, b in zip(want, got))
109
+ worst = max(worst, delta)
110
+ if int(np.argmax(want)) != int(np.argmax(got)) or delta > tolerance:
111
+ raise SystemExit(f"{onnx_path.name} disagrees with PyTorch on row {row['id']} (delta {delta:.4f})")
112
+ return worst
113
+
114
+
115
+ def main() -> None:
116
+ ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
117
+ ap.add_argument("--model", required=True, help="trained checkpoint directory")
118
+ ap.add_argument("--task", required=True)
119
+ ap.add_argument("--data", required=True, help="JSONL rows used to check the export (labels not needed)")
120
+ ap.add_argument("--max-tokens", type=int, default=512, help="serving token budget (match training)")
121
+ ap.add_argument("--fp32-only", action="store_true")
122
+ args = ap.parse_args()
123
+
124
+ task = load_task(args.task)
125
+ labels, question = task["labels"], task["questions"][0]
126
+ rows = load_rows(args.data, labels, require_labels=False)[:20]
127
+ agent = Agent(args.model, device="cpu")
128
+ agent.cfg["max_len"] = args.max_tokens
129
+
130
+ out_dir = Path(args.model) / "onnx"
131
+ out_dir.mkdir(exist_ok=True)
132
+ fp32 = out_dir / "model.onnx"
133
+ export_fp32(agent, row_state(rows[0]), question, fp32)
134
+ print(f"{fp32}: max probability difference vs PyTorch "
135
+ f"{check(agent, args.model, fp32, rows, question, labels, args.max_tokens, 1e-3):.5f}")
136
+ if args.fp32_only:
137
+ return
138
+
139
+ n_layers = json.loads((Path(args.model) / "encoder" / "config.json").read_text())["num_hidden_layers"]
140
+ nodes = encoder_weight_matmuls(fp32, n_layers)
141
+ if nodes is None:
142
+ print("int8: encoder is not ModernBERT/mmBERT-shaped; skipping block-wise quantization")
143
+ return
144
+ int8 = out_dir / "model-int8-blockwise.onnx"
145
+ quantize_blockwise(fp32, int8, nodes)
146
+ print(f"{int8}: max probability difference vs PyTorch "
147
+ f"{check(agent, args.model, int8, rows, question, labels, args.max_tokens, 0.05):.5f}")
148
+
149
+
150
+ if __name__ == "__main__":
151
+ main()
training/label.py ADDED
@@ -0,0 +1,135 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Label your prompts with two (or more) independent LLM annotators.
2
+
3
+ Raya's labels were made this way: two different strong models (Claude Opus and Claude Sonnet)
4
+ read the same written rubric and labelled every prompt without seeing each other's answers.
5
+ Where they agree you get a clean label; where they disagree the example keeps both votes and
6
+ train.py learns a 50/50 target instead of a wrong hard label.
7
+
8
+ export OPENAI_API_KEY=... # or any OpenAI-compatible provider
9
+ python label.py --task task.example.json --data prompts.jsonl --out labelled.jsonl \\
10
+ --annotator <model-a> \\
11
+ --annotator <model-b>@https://api.anthropic.com/v1/#ANTHROPIC_API_KEY
12
+
13
+ An annotator is ``MODEL[@BASE_URL][#API_KEY_ENV]``: BASE_URL defaults to
14
+ https://api.openai.com/v1/ and API_KEY_ENV to OPENAI_API_KEY. Use models from different
15
+ families so their mistakes are independent. Input rows need a ``prompt`` (or ``state``); any
16
+ existing labels are ignored. Re-running resumes: rows already labelled in --out are skipped.
17
+ """
18
+
19
+ from __future__ import annotations
20
+
21
+ import argparse
22
+ import json
23
+ import os
24
+ import re
25
+ import threading
26
+ import time
27
+ import urllib.error
28
+ import urllib.request
29
+ from concurrent.futures import ThreadPoolExecutor, as_completed
30
+ from pathlib import Path
31
+
32
+ from common import EXCLUDE, load_rows, load_task, row_state
33
+
34
+ DEFAULT_BASE = "https://api.openai.com/v1/"
35
+
36
+
37
+ def parse_annotator(spec: str) -> dict:
38
+ spec, _, key_env = spec.partition("#")
39
+ model, _, base = spec.partition("@")
40
+ return {"model": model, "base": (base or DEFAULT_BASE).rstrip("/") + "/",
41
+ "key_env": key_env or "OPENAI_API_KEY", "name": model}
42
+
43
+
44
+ def system_prompt(task: dict) -> str:
45
+ labels = task["labels"]
46
+ rubric = task.get("rubric") or "\n".join(
47
+ f"- {label}: {task['questions'][0]['criteria'][label]}" for label in labels
48
+ if isinstance(task["questions"][0].get("criteria"), dict))
49
+ return (
50
+ "You are an independent annotator building a training set.\n"
51
+ f"{task['questions'][0]['instructions']}\n\n{rubric}\n\n"
52
+ f'Use "{EXCLUDE}" only if the input is unusable (empty, gibberish, no discernible request).\n'
53
+ "Judge the input yourself; do not guess from keywords. "
54
+ f'Answer with JSON only: {{"label": one of {labels + [EXCLUDE]}}}'
55
+ )
56
+
57
+
58
+ def call(annotator: dict, system: str, user: str, labels: list[str], retries: int = 6) -> str:
59
+ key = os.environ.get(annotator["key_env"])
60
+ if not key:
61
+ raise SystemExit(f"set {annotator['key_env']} for annotator {annotator['model']}")
62
+ body = json.dumps({"model": annotator["model"], "temperature": 0, "max_tokens": 50,
63
+ "messages": [{"role": "system", "content": system}, {"role": "user", "content": user}]})
64
+ request = urllib.request.Request(annotator["base"] + "chat/completions", data=body.encode(), headers={
65
+ "Content-Type": "application/json", "Authorization": f"Bearer {key}", "x-api-key": key,
66
+ "anthropic-version": "2023-06-01"})
67
+ for attempt in range(retries):
68
+ try:
69
+ with urllib.request.urlopen(request, timeout=120) as response:
70
+ text = json.load(response)["choices"][0]["message"]["content"] or ""
71
+ match = re.search(r"\{.*?\}", text, re.S)
72
+ label = json.loads(match.group(0)).get("label") if match else None
73
+ if label in labels or label == EXCLUDE:
74
+ return label
75
+ raise ValueError(f"unexpected answer {text[:80]!r}")
76
+ except urllib.error.HTTPError as exc:
77
+ if exc.code not in (408, 409, 429) and exc.code < 500:
78
+ raise SystemExit(f"{annotator['model']}: HTTP {exc.code} {exc.read()[:200]!r}") from exc
79
+ error = exc
80
+ except (urllib.error.URLError, TimeoutError, ValueError, KeyError) as exc:
81
+ error = exc
82
+ time.sleep(min(60, 2 ** attempt))
83
+ raise RuntimeError(f"{annotator['model']} failed after {retries} attempts: {error}")
84
+
85
+
86
+ def main() -> None:
87
+ ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
88
+ ap.add_argument("--task", required=True)
89
+ ap.add_argument("--data", required=True, help="JSONL of rows with a 'prompt' or 'state'")
90
+ ap.add_argument("--out", required=True, help="labelled JSONL (appended; re-run to resume)")
91
+ ap.add_argument("--annotator", action="append", required=True, help="MODEL[@BASE_URL][#API_KEY_ENV]; repeat")
92
+ ap.add_argument("--workers", type=int, default=8)
93
+ ap.add_argument("--max-chars", type=int, default=6000, help="truncate long inputs sent to annotators")
94
+ args = ap.parse_args()
95
+
96
+ task = load_task(args.task)
97
+ labels = task["labels"]
98
+ annotators = [parse_annotator(a) for a in args.annotator]
99
+ if len(annotators) < 2:
100
+ print("warning: one annotator gives hard labels only; two different models are recommended")
101
+ system = system_prompt(task)
102
+ rows = load_rows(args.data, labels, require_labels=False)
103
+ out = Path(args.out)
104
+ done = {json.loads(l)["id"] for l in out.read_text(encoding="utf-8").splitlines() if l.strip()} if out.exists() else set()
105
+ todo = [r for r in rows if r["id"] not in done]
106
+ print(f"{len(rows)} rows, {len(done)} already labelled, {len(todo)} to go, {len(annotators)} annotators")
107
+
108
+ lock = threading.Lock()
109
+
110
+ def label_row(row: dict) -> dict:
111
+ state = row_state(row)
112
+ user = state["prompt"] if isinstance(state, dict) and set(state) == {"prompt"} else json.dumps(state, ensure_ascii=False)
113
+ if len(user) > args.max_chars:
114
+ user = user[:args.max_chars] + " …[truncated]"
115
+ votes = {a["name"]: call(a, system, user, labels) for a in annotators}
116
+ clean = {k: v for k, v in row.items() if k not in ("label", "labels", "target", "gold")}
117
+ return dict(clean, labels=list(votes.values()), annotators=votes)
118
+
119
+ agree = labelled = 0
120
+ with ThreadPoolExecutor(args.workers) as pool, out.open("a", encoding="utf-8") as sink:
121
+ for future in as_completed(pool.submit(label_row, r) for r in todo):
122
+ result = future.result()
123
+ with lock:
124
+ sink.write(json.dumps(result, ensure_ascii=False) + "\n")
125
+ sink.flush()
126
+ labelled += 1
127
+ agree += len(set(result["labels"])) == 1
128
+ if labelled % 50 == 0:
129
+ print(f"{labelled}/{len(todo)} labelled, annotators agree on {agree / labelled:.0%}", flush=True)
130
+ if labelled:
131
+ print(f"done: {labelled} labelled, annotators agree on {agree / labelled:.0%}")
132
+
133
+
134
+ if __name__ == "__main__":
135
+ main()
training/requirements-onnx.txt ADDED
@@ -0,0 +1,5 @@
 
 
 
 
 
 
1
+ # Only needed for export_onnx.py and CPU serving. Tested: onnx 1.23.0, onnxscript 0.7.2, onnxruntime 1.30.0.
2
+ -r requirements.txt
3
+ onnx>=1.17
4
+ onnxscript>=0.2
5
+ onnxruntime>=1.20
training/requirements.txt ADDED
@@ -0,0 +1,8 @@
 
 
 
 
 
 
 
 
 
1
+ # Tested with these exact versions (September 2026). laya is pinned because its tokenisation
2
+ # and temperature handling decide the model's inputs; the rest are minimums.
3
+ laya==0.3.20
4
+ torch>=2.4
5
+ transformers>=4.48
6
+ safetensors>=0.4
7
+ numpy>=1.26
8
+ huggingface_hub>=0.26
training/task.example.json ADDED
@@ -0,0 +1,34 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "description": "3-tier LLM routing, the task Raya was trained on. Copy this file and edit it for your own task.",
3
+ "labels": ["small_model", "medium_model", "frontier_model"],
4
+ "questions": [
5
+ {
6
+ "type": "choice",
7
+ "instructions": "Route this prompt to a model.",
8
+ "criteria": {
9
+ "small_model": "simple requests",
10
+ "medium_model": "moderately complex requests",
11
+ "frontier_model": "very hard requests"
12
+ }
13
+ },
14
+ {
15
+ "type": "choice",
16
+ "instructions": "Which model tier should answer this user prompt? Pick the cheapest tier that will still answer it well.",
17
+ "criteria": {
18
+ "small_model": "greetings, trivial facts, single-sentence translation or rewrite, typo fixes, simple arithmetic, one-liner code questions, very short simple creative requests",
19
+ "medium_model": "normal multi-paragraph writing, emails, essays, standard coding tasks, summarizing/translating/rewriting a provided text, explaining well-known concepts, routine analysis",
20
+ "frontier_model": "hard multi-step reasoning, non-trivial math or proofs, complex system design or large tricky code, expert legal/financial/medical/scientific analysis, research-grade or strategy work"
21
+ }
22
+ },
23
+ {
24
+ "type": "score",
25
+ "instructions": "How difficult is this prompt for an AI assistant to answer well?",
26
+ "criteria": [
27
+ "simple: a small fast model answers it perfectly",
28
+ "moderate: needs a capable general model",
29
+ "hard: needs the strongest frontier model"
30
+ ]
31
+ }
32
+ ],
33
+ "rubric": "Decide which model tier should answer the prompt for a production assistant that wants the CHEAPEST model that will still answer well.\n- small_model: greetings, trivial facts, single-sentence translation/rewrite/typo fix, simple arithmetic, one-liner code questions, very short simple creative requests.\n- medium_model: normal multi-paragraph writing, emails, essays, standard coding tasks, summarizing/translating/rewriting a provided text, explanations of well-known concepts, routine analysis.\n- frontier_model: hard multi-step reasoning, non-trivial math/proofs, complex system design or large/tricky code, expert legal/financial/medical/scientific analysis, long research-grade or strategy work synthesizing many sources.\nJudge the prompt in its own language. Do not aim for class balance."
34
+ }
training/train.py ADDED
@@ -0,0 +1,354 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Fine-tune a Laya decision model (a "System-1" model) on your own labelled data.
2
+
3
+ This is the recipe behind TextCortex/raya, generalised to any choice/score task:
4
+
5
+ * soft targets: each label's share of the annotator votes (a 50/50 target where two disagree);
6
+ * every question phrasing in the task is trained on, and choice options are shuffled, so the
7
+ model learns the content rather than the wording or the option order;
8
+ * soft-target cross-entropy with 1/sqrt(frequency) class weights, so rare labels still count;
9
+ * the best epoch is picked on validation, then a temperature per question is fitted there so
10
+ the output probabilities are calibrated;
11
+ * the result is a normal Laya checkpoint: ``laya.Agent("<out>")`` loads it.
12
+
13
+ Quick start (CPU/Apple Silicon works for small data; a GPU is much faster):
14
+
15
+ python train.py --task task.example.json --data data/example.jsonl --out my-router
16
+
17
+ Start from Raya instead of stock Laya to adapt the router to your own traffic:
18
+
19
+ python train.py --task task.example.json --data my_data.jsonl --base TextCortex/raya --out my-router
20
+ """
21
+
22
+ from __future__ import annotations
23
+
24
+ import argparse
25
+ import json
26
+ import math
27
+ import os
28
+ import random
29
+ import shutil
30
+ import time
31
+ from pathlib import Path
32
+
33
+ os.environ.setdefault("USE_TF", "0")
34
+
35
+ import numpy as np
36
+ import torch
37
+ import torch.nn.functional as F
38
+ from safetensors.torch import save_file
39
+
40
+ import laya
41
+ from laya.common import QTYPES, build_sequence, collate_items, temp_bucket
42
+
43
+ from common import load_rows, load_task, row_state
44
+
45
+ BASE_FILES = ["rl_agent_config.json", "model.safetensors", "tokenizer/*", "encoder/*"]
46
+
47
+
48
+ def resolve_base(base: str, subfolder: str | None) -> Path:
49
+ """Local directory holding the base checkpoint (downloads it from the Hub if needed)."""
50
+ if Path(base).is_dir():
51
+ return Path(base) / subfolder if subfolder else Path(base)
52
+ from huggingface_hub import snapshot_download
53
+
54
+ patterns = [f"{subfolder}/{p}" for p in BASE_FILES] if subfolder else BASE_FILES
55
+ local = Path(snapshot_download(base, allow_patterns=patterns))
56
+ return local / subfolder if subfolder else local
57
+
58
+
59
+ def internal(question: dict) -> dict:
60
+ return laya.Agent._to_internal(question)
61
+
62
+
63
+ def make_item(tok, cfg, state, qi: int, question: dict, labels: list[str], target: list[float], shuffle: bool):
64
+ """Tokenise one (example, question) pair and align its target with the option order."""
65
+ q = internal(question)
66
+ if q["t"] == "choice":
67
+ keys = list(q["crit"].keys())
68
+ by_label = dict(zip(labels, target))
69
+ order = list(range(len(keys)))
70
+ if shuffle:
71
+ random.shuffle(order)
72
+ seq, markers = build_sequence(tok, state, q, cfg["max_len"], cfg["head_max_len"], option_order=order)
73
+ tgt = [by_label[keys[j]] for j in order]
74
+ else: # ordinal score: the option order is the meaning, never shuffle
75
+ seq, markers = build_sequence(tok, state, q, cfg["max_len"], cfg["head_max_len"])
76
+ tgt = list(target)
77
+ return {"ids": seq, "markers": markers, "qtype": QTYPES[q["t"]], "target": tgt,
78
+ "question": qi, "label_target": list(target)}
79
+
80
+
81
+ def length_batches(items, bs: int, shuffle: bool):
82
+ """Batches of similar length (less padding), in random order when training."""
83
+ idx = sorted(range(len(items)), key=lambda i: len(items[i]["ids"]))
84
+ chunks = [idx[i:i + bs] for i in range(0, len(idx), bs)]
85
+ if shuffle:
86
+ random.shuffle(chunks)
87
+ return chunks
88
+
89
+
90
+ def forward(model, batch, dev):
91
+ with torch.autocast(device_type=dev.type, dtype=torch.bfloat16, enabled=dev.type == "cuda"):
92
+ logits, _ = model(batch["input_ids"].to(dev), batch["attention_mask"].to(dev), batch["marker_pos"].to(dev),
93
+ batch["marker_mask"].to(dev), batch["qtype"].to(dev))
94
+ return logits.float()
95
+
96
+
97
+ @torch.no_grad()
98
+ def predict_logits(model, items, pad_id, dev, bs=16):
99
+ model.eval()
100
+ out = [None] * len(items)
101
+ for chunk in length_batches(items, bs, shuffle=False):
102
+ batch = collate_items([[items[i] for i in chunk]], pad_id)
103
+ logits = forward(model, batch, dev).cpu()
104
+ for j, i in enumerate(chunk):
105
+ out[i] = logits[j, :len(items[i]["markers"])].numpy()
106
+ return out
107
+
108
+
109
+ def softmax(z: np.ndarray, t: float = 1.0) -> np.ndarray:
110
+ z = z / t
111
+ p = np.exp(z - z.max())
112
+ return p / p.sum()
113
+
114
+
115
+ def evaluate(logits, items, rows_of_items, n_questions: int, n_labels: int, temps=None):
116
+ """Accuracy and macro-F1 per question, on examples with a single gold label."""
117
+ res = {}
118
+ for qi in range(n_questions):
119
+ preds, golds = [], []
120
+ for lg, it, row in zip(logits, items, rows_of_items):
121
+ if it["question"] != qi or row["gold"] is None:
122
+ continue
123
+ preds.append(int(softmax(lg, (temps or {}).get(qi, 1.0)).argmax()))
124
+ golds.append(int(np.argmax(it["target"])))
125
+ if not golds:
126
+ continue
127
+ f1s = []
128
+ for k in range(n_labels):
129
+ tp = sum(p == k == g for p, g in zip(preds, golds))
130
+ fp = sum(p == k != g for p, g in zip(preds, golds))
131
+ fn = sum(g == k != p for p, g in zip(preds, golds))
132
+ f1s.append(2 * tp / (2 * tp + fp + fn) if tp else 0.0)
133
+ res[qi] = {"acc": round(float(np.mean([p == g for p, g in zip(preds, golds)])), 4),
134
+ "macro_f1": round(float(np.mean(f1s)), 4), "n": len(golds)}
135
+ return res
136
+
137
+
138
+ def soft_nll(logits, items) -> float:
139
+ total = 0.0
140
+ for lg, it in zip(logits, items):
141
+ total -= float((np.array(it["target"]) * np.log(softmax(lg) + 1e-12)).sum())
142
+ return total / len(items)
143
+
144
+
145
+ def fit_temperature(logits, items, qi: int) -> float:
146
+ """Temperature minimising the soft-target NLL of one question on validation."""
147
+ sel = [(lg, np.array(it["target"])) for lg, it in zip(logits, items) if it["question"] == qi]
148
+ best_t, best_nll = 1.0, float("inf")
149
+ for t in np.arange(0.5, 5.01, 0.05):
150
+ nll = -sum(float((tg * np.log(softmax(lg, t) + 1e-12)).sum()) for lg, tg in sel)
151
+ if nll < best_nll:
152
+ best_t, best_nll = round(float(t), 2), nll
153
+ return best_t
154
+
155
+
156
+ def split_rows(rows, val_frac: float, seed: int):
157
+ """Rows marked ``"split": "val"`` are validation; otherwise a random ``val_frac`` share is."""
158
+ marked = [r for r in rows if r.get("split") == "val"]
159
+ if marked:
160
+ return [r for r in rows if r.get("split") != "val"], marked
161
+ rows = list(rows)
162
+ random.Random(seed).shuffle(rows)
163
+ n_val = max(1, int(len(rows) * val_frac))
164
+ return rows[n_val:], rows[:n_val]
165
+
166
+
167
+ def main() -> None:
168
+ ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
169
+ ap.add_argument("--task", required=True, help="task.json: labels + question phrasings")
170
+ ap.add_argument("--data", required=True, help="training JSONL (see common.py for the format)")
171
+ ap.add_argument("--val-data", help="optional separate validation JSONL")
172
+ ap.add_argument("--out", required=True, help="output checkpoint directory")
173
+ ap.add_argument("--base", default="convaiinnovations/laya",
174
+ help="base checkpoint: a Hub repo id or local dir (e.g. TextCortex/raya to adapt Raya)")
175
+ ap.add_argument("--subfolder", default=None,
176
+ help="checkpoint subfolder; defaults to 'multilingual' for convaiinnovations/laya "
177
+ "(pass --subfolder . for Laya's English ModernBERT-large checkpoint)")
178
+ ap.add_argument("--encoder", default=None,
179
+ help="build a fresh decision head on this HF encoder (e.g. jhu-clsp/mmBERT-small) instead of --base")
180
+ ap.add_argument("--epochs", type=int, default=3)
181
+ ap.add_argument("--batch-size", type=int, default=16)
182
+ ap.add_argument("--lr-encoder", type=float, default=2e-5)
183
+ ap.add_argument("--lr-head", type=float, default=None, help="default 1e-4, or 3e-4 for a fresh head")
184
+ ap.add_argument("--max-tokens", type=int, default=512, help="input token budget while training")
185
+ ap.add_argument("--val-frac", type=float, default=0.1)
186
+ ap.add_argument("--select", choices=["acc", "nll"], default="acc", help="best-epoch criterion on validation")
187
+ ap.add_argument("--seed", type=int, default=0)
188
+ ap.add_argument("--device", default=None, help="cuda | mps | cpu (default: best available)")
189
+ ap.add_argument("--train-embeddings", action="store_true",
190
+ help="also train the token-embedding table (frozen by default: saves memory, rarely helps)")
191
+ ap.add_argument("--push-to-hub", metavar="REPO_ID", help="upload the result to this Hugging Face model repo")
192
+ ap.add_argument("--private", action="store_true", help="with --push-to-hub: create the repo as private")
193
+ args = ap.parse_args()
194
+
195
+ random.seed(args.seed)
196
+ np.random.seed(args.seed)
197
+ torch.manual_seed(args.seed)
198
+
199
+ task = load_task(args.task)
200
+ labels, questions = task["labels"], task["questions"]
201
+ rows = load_rows(args.data, labels)
202
+ if args.val_data:
203
+ train_rows, val_rows = rows, load_rows(args.val_data, labels)
204
+ else:
205
+ train_rows, val_rows = split_rows(rows, args.val_frac, args.seed)
206
+ mass = np.array([sum(r["target"][i] for r in train_rows) for i in range(len(labels))])
207
+ print(f"train {len(train_rows)} rows, validation {len(val_rows)} rows, {len(questions)} question phrasing(s)")
208
+ print("train label mass:", dict(zip(labels, mass.round(1).tolist())))
209
+ if (mass == 0).any():
210
+ raise SystemExit(f"no training examples for: {[l for l, m in zip(labels, mass) if m == 0]}")
211
+
212
+ dev = torch.device(args.device or ("cuda" if torch.cuda.is_available()
213
+ else "mps" if torch.backends.mps.is_available() else "cpu"))
214
+ if args.encoder:
215
+ from transformers import AutoTokenizer
216
+ from laya.common import build_model
217
+
218
+ # Borrow Laya's head/config layout, but train the decision head from scratch on a new encoder.
219
+ ref_dir = resolve_base("convaiinnovations/laya", "multilingual")
220
+ cfg = dict(json.loads((ref_dir / "rl_agent_config.json").read_text()), encoder=args.encoder,
221
+ max_len=args.max_tokens, head_max_len=192, temperature=[1.0, 1.0, 1.0], temperature_by_options={})
222
+ tok = AutoTokenizer.from_pretrained(args.encoder)
223
+ model = build_model(cfg, pretrained=True).to(dev)
224
+ base_dir = None
225
+ lr_head = args.lr_head or 3e-4
226
+ print(f"device {dev}: fresh decision head on {args.encoder}")
227
+ else:
228
+ subfolder = args.subfolder if args.subfolder is not None else (
229
+ "multilingual" if args.base == "convaiinnovations/laya" else None)
230
+ subfolder = None if subfolder in ("", ".") else subfolder
231
+ base_dir = resolve_base(args.base, subfolder)
232
+ agent = laya.Agent(str(base_dir), device=str(dev))
233
+ model, tok, cfg = agent.model, agent.tok, agent.cfg
234
+ lr_head = args.lr_head or 1e-4
235
+ print(f"device {dev}: fine-tuning {args.base}{'/' + subfolder if subfolder else ''}")
236
+
237
+ # Weight each example by its target's class weight (~1/sqrt(frequency), mean 1).
238
+ cw = (mass.sum() / mass) ** 0.5
239
+ cw = torch.tensor(cw / cw.mean(), dtype=torch.float32)
240
+ print("class weights:", {l: round(float(w), 2) for l, w in zip(labels, cw)})
241
+
242
+ if not args.train_embeddings:
243
+ for p in model.encoder.get_input_embeddings().parameters():
244
+ p.requires_grad_(False)
245
+ if dev.type != "cuda":
246
+ model.encoder.gradient_checkpointing_enable() # trade speed for memory off-GPU
247
+ train_cfg = dict(cfg, max_len=min(cfg["max_len"], args.max_tokens))
248
+
249
+ def items_for(rs, shuffle, c):
250
+ its, owners = [], []
251
+ for r in rs:
252
+ state = row_state(r)
253
+ for qi, q in enumerate(questions):
254
+ its.append(make_item(tok, c, state, qi, q, labels, r["target"], shuffle))
255
+ owners.append(r)
256
+ return its, owners
257
+
258
+ val_items, val_owners = items_for(val_rows, False, train_cfg)
259
+ before = predict_logits(model, val_items, tok.pad_token_id, dev)
260
+ print("validation before training:", json.dumps(evaluate(before, val_items, val_owners, len(questions), len(labels))))
261
+
262
+ enc_params = [p for p in model.encoder.parameters() if p.requires_grad]
263
+ head_params = [p for n, p in model.named_parameters() if not n.startswith("encoder.") and p.requires_grad]
264
+ opt = torch.optim.AdamW([{"params": enc_params, "lr": args.lr_encoder}, {"params": head_params, "lr": lr_head}],
265
+ weight_decay=0.01)
266
+ steps_per_epoch = math.ceil(len(train_rows) * len(questions) / args.batch_size)
267
+ total = steps_per_epoch * args.epochs
268
+ warm = max(1, int(0.1 * total))
269
+ sched = torch.optim.lr_scheduler.LambdaLR(
270
+ opt, lambda s: min(1.0, (s + 1) / warm) * max(0.0, (total - s) / max(1, total - warm)))
271
+
272
+ best_crit, best_state, best_logits = float("inf"), None, None
273
+ history, step, t0 = [], 0, time.time()
274
+ for epoch in range(args.epochs):
275
+ model.train()
276
+ items, _ = items_for(train_rows, True, train_cfg) # fresh option shuffle every epoch
277
+ running = []
278
+ for chunk in length_batches(items, args.batch_size, shuffle=True):
279
+ batch = collate_items([[items[i] for i in chunk]], tok.pad_token_id)
280
+ logits = forward(model, batch, dev)
281
+ tgt = batch["target"][:, :logits.size(1)].to(dev)
282
+ logp = F.log_softmax(logits.masked_fill(~batch["marker_mask"].to(dev), -1e4), -1)
283
+ per_example = -(tgt * logp).sum(-1)
284
+ w = torch.tensor([float((torch.tensor(items[i]["label_target"]) * cw).sum()) for i in chunk], device=dev)
285
+ loss = (per_example * w).sum() / w.sum()
286
+ opt.zero_grad()
287
+ loss.backward()
288
+ torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
289
+ opt.step()
290
+ sched.step()
291
+ step += 1
292
+ running.append(loss.item())
293
+ if dev.type == "mps" and step % 25 == 0:
294
+ torch.mps.empty_cache()
295
+ if step % 50 == 0:
296
+ print(f"epoch {epoch} step {step}/{total} loss {np.mean(running[-50:]):.4f} ({time.time() - t0:.0f}s)",
297
+ flush=True)
298
+ val_logits = predict_logits(model, val_items, tok.pad_token_id, dev)
299
+ ev = evaluate(val_logits, val_items, val_owners, len(questions), len(labels))
300
+ mean_acc = float(np.mean([v["acc"] for v in ev.values()])) if ev else 0.0
301
+ nll = soft_nll(val_logits, val_items)
302
+ crit = -mean_acc if args.select == "acc" else nll
303
+ history.append({"epoch": epoch, "val_mean_acc": round(mean_acc, 4), "val_soft_nll": round(nll, 4),
304
+ "val": ev, "seconds": round(time.time() - t0)})
305
+ print(f"epoch {epoch}: validation {json.dumps(ev)} mean acc {mean_acc:.4f} soft NLL {nll:.4f}", flush=True)
306
+ if crit < best_crit:
307
+ best_crit = crit
308
+ best_state = {k: v.detach().cpu().clone() for k, v in model.state_dict().items()}
309
+ best_logits = val_logits
310
+
311
+ temps = {qi: fit_temperature(best_logits, val_items, qi) for qi in range(len(questions))}
312
+ final = evaluate(best_logits, val_items, val_owners, len(questions), len(labels), temps)
313
+ print("fitted temperatures:", temps, "validation after calibration:", json.dumps(final))
314
+
315
+ out = Path(args.out)
316
+ if out.exists():
317
+ shutil.rmtree(out)
318
+ if base_dir is None:
319
+ out.mkdir(parents=True)
320
+ tok.save_pretrained(out / "tokenizer")
321
+ model.encoder.config.save_pretrained(out / "encoder")
322
+ else:
323
+ shutil.copytree(base_dir, out, ignore=shutil.ignore_patterns(
324
+ "model.safetensors", "*.onnx", "onnx", "multilingual", "typed-decisions", "assets", "*.md", ".git*", ".cache"))
325
+ save_file({k: v.contiguous() for k, v in best_state.items()}, str(out / "model.safetensors"))
326
+
327
+ # Laya applies one temperature per (question type, option count) bucket; average within a bucket.
328
+ buckets: dict[str, list[float]] = {}
329
+ for qi, q in enumerate(questions):
330
+ iq = internal(q)
331
+ buckets.setdefault(temp_bucket(QTYPES[iq["t"]], len(iq["crit"])), []).append(temps[qi])
332
+ new_cfg = dict(cfg)
333
+ new_cfg["temperature_by_options"] = {**cfg.get("temperature_by_options", {}),
334
+ **{b: round(float(np.mean(ts)), 2) for b, ts in buckets.items()}}
335
+ new_cfg["training"] = dict(cfg.get("training", {}), fine_tuned={
336
+ "base": args.encoder or args.base, "train_rows": len(train_rows), "val_rows": len(val_rows),
337
+ "epochs": args.epochs, "select": args.select, "seed": args.seed, "labels": labels})
338
+ (out / "rl_agent_config.json").write_text(json.dumps(new_cfg, indent=2))
339
+ (out / "task.json").write_text(json.dumps(task, indent=2, ensure_ascii=False))
340
+ (out / "training_log.json").write_text(json.dumps(
341
+ {"args": vars(args), "history": history, "temperatures": temps, "validation": final}, indent=2))
342
+ print(f"saved {out} (load it with laya.Agent({str(out)!r}))")
343
+
344
+ if args.push_to_hub:
345
+ from huggingface_hub import HfApi
346
+
347
+ api = HfApi()
348
+ api.create_repo(args.push_to_hub, private=args.private, exist_ok=True)
349
+ api.upload_folder(folder_path=str(out), repo_id=args.push_to_hub)
350
+ print(f"uploaded to https://huggingface.co/{args.push_to_hub}")
351
+
352
+
353
+ if __name__ == "__main__":
354
+ main()