Add training kit: train your own System-1 model
Browse filesLabel (two independent LLM annotators), train, evaluate and export (ONNX fp32 + block-wise int8) scripts generalised from Raya's training. Model files unchanged.
- README.md +18 -0
- training/README.md +215 -0
- training/common.py +125 -0
- training/data/example.jsonl +44 -0
- training/data/example_test.jsonl +13 -0
- training/data/unlabelled.jsonl +4 -0
- training/evaluate.py +97 -0
- training/export_onnx.py +151 -0
- training/label.py +135 -0
- training/requirements-onnx.txt +5 -0
- training/requirements.txt +8 -0
- training/task.example.json +34 -0
- training/train.py +354 -0
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()
|