Text Classification
Transformers
Safetensors
laya
system-one
calibrated-decisions
rlcd
classification
routing
scoring
guardrails
moderation
reinforcement-learning
commercial-use
Instructions to use sekkit/laya with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use sekkit/laya with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-classification", model="sekkit/laya")# Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("sekkit/laya", device_map="auto") - Notebooks
- Google Colab
- Kaggle
Commit ·
97f34fe
0
Parent(s):
Duplicate from convaiinnovations/laya
Browse filesCo-authored-by: Convai Innovations <convaiinnovations@users.noreply.huggingface.co>
- .gitattributes +41 -0
- README.md +266 -0
- assets/laya_benchmark.png +3 -0
- assets/laya_benchmark_common.png +3 -0
- assets/laya_vs_jev.png +3 -0
- assets/laya_vs_jev_full.png +3 -0
- assets/logo-lockup-dark.png +0 -0
- assets/logo-lockup-dark.svg +22 -0
- assets/logo-lockup.png +0 -0
- assets/logo-lockup.svg +22 -0
- assets/logo-mark-ink.png +0 -0
- assets/logo-mark-ink.svg +20 -0
- assets/logo-mark-mono.svg +20 -0
- assets/logo-mark.png +0 -0
- assets/logo-mark.svg +20 -0
- email_utils.py +71 -0
- encoder/config.json +84 -0
- eval/benchmark_comparison.png +3 -0
- eval/reliability_eval_in.png +0 -0
- eval/reliability_eval_zs.png +0 -0
- eval/results.json +142 -0
- eval/results.md +53 -0
- model.safetensors +3 -0
- multilingual/encoder/config.json +79 -0
- multilingual/model.safetensors +3 -0
- multilingual/rl_agent_config.json +26 -0
- multilingual/tokenizer/tokenizer.json +3 -0
- multilingual/tokenizer/tokenizer_config.json +22 -0
- rl_agent_api.py +78 -0
- rl_agent_config.json +33 -0
- rl_common.py +408 -0
- tokenizer/tokenizer.json +0 -0
- tokenizer/tokenizer_config.json +14 -0
- typed-decisions/encoder/config.json +84 -0
- typed-decisions/model.safetensors +3 -0
- typed-decisions/rl_agent_config.json +36 -0
- typed-decisions/tokenizer/tokenizer.json +0 -0
- typed-decisions/tokenizer/tokenizer_config.json +15 -0
.gitattributes
ADDED
|
@@ -0,0 +1,41 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
*.7z filter=lfs diff=lfs merge=lfs -text
|
| 2 |
+
*.arrow filter=lfs diff=lfs merge=lfs -text
|
| 3 |
+
*.bin filter=lfs diff=lfs merge=lfs -text
|
| 4 |
+
*.bz2 filter=lfs diff=lfs merge=lfs -text
|
| 5 |
+
*.ckpt filter=lfs diff=lfs merge=lfs -text
|
| 6 |
+
*.ftz filter=lfs diff=lfs merge=lfs -text
|
| 7 |
+
*.gz filter=lfs diff=lfs merge=lfs -text
|
| 8 |
+
*.h5 filter=lfs diff=lfs merge=lfs -text
|
| 9 |
+
*.joblib filter=lfs diff=lfs merge=lfs -text
|
| 10 |
+
*.lfs.* filter=lfs diff=lfs merge=lfs -text
|
| 11 |
+
*.mlmodel filter=lfs diff=lfs merge=lfs -text
|
| 12 |
+
*.model filter=lfs diff=lfs merge=lfs -text
|
| 13 |
+
*.msgpack filter=lfs diff=lfs merge=lfs -text
|
| 14 |
+
*.npy filter=lfs diff=lfs merge=lfs -text
|
| 15 |
+
*.npz filter=lfs diff=lfs merge=lfs -text
|
| 16 |
+
*.onnx filter=lfs diff=lfs merge=lfs -text
|
| 17 |
+
*.ot filter=lfs diff=lfs merge=lfs -text
|
| 18 |
+
*.parquet filter=lfs diff=lfs merge=lfs -text
|
| 19 |
+
*.pb filter=lfs diff=lfs merge=lfs -text
|
| 20 |
+
*.pickle filter=lfs diff=lfs merge=lfs -text
|
| 21 |
+
*.pkl filter=lfs diff=lfs merge=lfs -text
|
| 22 |
+
*.pt filter=lfs diff=lfs merge=lfs -text
|
| 23 |
+
*.pth filter=lfs diff=lfs merge=lfs -text
|
| 24 |
+
*.rar filter=lfs diff=lfs merge=lfs -text
|
| 25 |
+
*.safetensors filter=lfs diff=lfs merge=lfs -text
|
| 26 |
+
saved_model/**/* filter=lfs diff=lfs merge=lfs -text
|
| 27 |
+
*.tar.* filter=lfs diff=lfs merge=lfs -text
|
| 28 |
+
*.tar filter=lfs diff=lfs merge=lfs -text
|
| 29 |
+
*.tflite filter=lfs diff=lfs merge=lfs -text
|
| 30 |
+
*.tgz filter=lfs diff=lfs merge=lfs -text
|
| 31 |
+
*.wasm filter=lfs diff=lfs merge=lfs -text
|
| 32 |
+
*.xz filter=lfs diff=lfs merge=lfs -text
|
| 33 |
+
*.zip filter=lfs diff=lfs merge=lfs -text
|
| 34 |
+
*.zst filter=lfs diff=lfs merge=lfs -text
|
| 35 |
+
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
| 36 |
+
eval/benchmark_comparison.png filter=lfs diff=lfs merge=lfs -text
|
| 37 |
+
multilingual/tokenizer/tokenizer.json filter=lfs diff=lfs merge=lfs -text
|
| 38 |
+
assets/laya_vs_jev.png filter=lfs diff=lfs merge=lfs -text
|
| 39 |
+
assets/laya_benchmark.png filter=lfs diff=lfs merge=lfs -text
|
| 40 |
+
assets/laya_benchmark_common.png filter=lfs diff=lfs merge=lfs -text
|
| 41 |
+
assets/laya_vs_jev_full.png filter=lfs diff=lfs merge=lfs -text
|
README.md
ADDED
|
@@ -0,0 +1,266 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
---
|
| 2 |
+
license: apache-2.0
|
| 3 |
+
library_name: transformers
|
| 4 |
+
pipeline_tag: text-classification
|
| 5 |
+
tags: [laya, system-one, calibrated-decisions, rlcd, classification, routing, scoring, guardrails, moderation, reinforcement-learning, commercial-use]
|
| 6 |
+
---
|
| 7 |
+
|
| 8 |
+
# Laya
|
| 9 |
+
|
| 10 |
+
**Multilingual, non-autoregressive System 1 decision model.** Give it a **state** (text, email, ticket, or JSON) and **typed questions**; it returns typed answers with mathematically calibrated probabilities in a single forward pass (~33 ms) across 100+ languages. Trained with reinforcement learning against strictly proper scoring rules (**RLCD**), so reporting honest probabilities is the only way to maximise reward. It never generates text, so there is nothing to parse and nothing to hallucinate.
|
| 11 |
+
|
| 12 |
+
<p align="center">
|
| 13 |
+
<img src="https://raw.githubusercontent.com/NandhaKishorM/laya/main/assets/laya_vs_jev_full.png" alt="Laya versus TypeSafe Jev: accuracy, every application workflow, all 51 languages, speed, calibration and routing cost" width="100%" />
|
| 14 |
+
</p>
|
| 15 |
+
|
| 16 |
+
**This repo holds all three checkpoints** and is the hub for the family. The English checkpoint is at the repo root; the other two are bundled subfolders, and only the one you request is downloaded:
|
| 17 |
+
|
| 18 |
+
| Checkpoint | Backbone Encoder | Params | Context | Best at |
|
| 19 |
+
|---|---|---|---|---|
|
| 20 |
+
| **`convaiinnovations/laya`** (this repo root) | ModernBERT-large | 421M | 512 | English text, guardrails, email triage |
|
| 21 |
+
| [`convaiinnovations/laya-multilingual`](https://huggingface.co/convaiinnovations/laya-multilingual) | mmBERT-base | 322M | 1024 (up to 8k) | 100+ languages, ~2.2x faster |
|
| 22 |
+
| [`convaiinnovations/laya-typed-decisions`](https://huggingface.co/convaiinnovations/laya-typed-decisions) | ModernBERT-large | 421M | 1024 | the four typed-decisions workflows (0.766 acc) |
|
| 23 |
+
|
| 24 |
+
---
|
| 25 |
+
|
| 26 |
+
## Quickstart: Route Mode (Recommended)
|
| 27 |
+
|
| 28 |
+
Laya's built-in **`Router`** is the recommended way to use Laya in production. It evaluates any state in any language, automatically detects scripts and languages in sub-milliseconds, and dispatches to the optimal checkpoint in a single forward pass.
|
| 29 |
+
|
| 30 |
+
```bash
|
| 31 |
+
pip install laya
|
| 32 |
+
```
|
| 33 |
+
|
| 34 |
+
```python
|
| 35 |
+
import laya
|
| 36 |
+
from laya import Router
|
| 37 |
+
|
| 38 |
+
# Preload checkpoints into memory for instant sub-35ms routing
|
| 39 |
+
router = Router(preload=True)
|
| 40 |
+
|
| 41 |
+
state = {
|
| 42 |
+
"from": "user@acme.com",
|
| 43 |
+
"subject": "Duplicate charge on invoice #4411",
|
| 44 |
+
"body": "Hi, we were billed twice for March. Please refund the duplicate today or we will cancel our plan."
|
| 45 |
+
}
|
| 46 |
+
|
| 47 |
+
questions = {
|
| 48 |
+
"department": {
|
| 49 |
+
"type": "choice",
|
| 50 |
+
"instructions": "Which department should handle this request?",
|
| 51 |
+
"criteria": {
|
| 52 |
+
"billing": "invoices, payments, refunds",
|
| 53 |
+
"technical": "bugs, outages, system errors",
|
| 54 |
+
"sales": "pricing, new contracts",
|
| 55 |
+
"other": "everything else"
|
| 56 |
+
}
|
| 57 |
+
},
|
| 58 |
+
"urgency": {
|
| 59 |
+
"type": "score",
|
| 60 |
+
"instructions": "How urgent is this request?",
|
| 61 |
+
"criteria": ["not urgent", "soon", "critical deadline or blocking issue"]
|
| 62 |
+
},
|
| 63 |
+
"churn_risk": {
|
| 64 |
+
"type": "noul",
|
| 65 |
+
"instructions": "Does the user threaten to cancel or leave?"
|
| 66 |
+
},
|
| 67 |
+
"refund_requested": {
|
| 68 |
+
"type": "noul",
|
| 69 |
+
"instructions": "Does the user explicitly request a refund?"
|
| 70 |
+
}
|
| 71 |
+
}
|
| 72 |
+
|
| 73 |
+
# 1. English state -> automatically routed to ModernBERT-large (39.5 ms)
|
| 74 |
+
res_en = router.predict(state, questions)
|
| 75 |
+
print("Department :", res_en["answers"]["department"]["choice"]) # -> billing (confidence: 0.94)
|
| 76 |
+
print("Routing :", res_en["routing"]["model"]) # -> english
|
| 77 |
+
|
| 78 |
+
# 2. Hindi state -> automatically routed to mmBERT-base (100+ languages, 32.8 ms)
|
| 79 |
+
res_hi = router.predict({"body": "मुझसे दो बार शुल्क लिया गया, कृपया पैसे वापस करें।"}, questions)
|
| 80 |
+
print("Department :", res_hi["answers"]["department"]["choice"]) # -> billing (confidence: 0.86)
|
| 81 |
+
print("Routing :", res_hi["routing"]["model"]) # -> multilingual
|
| 82 |
+
|
| 83 |
+
# 3. Explicit override when you already know the checkpoint
|
| 84 |
+
res_td = router.predict(state, questions, model="typed-decisions")
|
| 85 |
+
```
|
| 86 |
+
|
| 87 |
+
Every result carries full routing metadata explaining why the choice was made:
|
| 88 |
+
|
| 89 |
+
```python
|
| 90 |
+
res_hi["routing"]
|
| 91 |
+
# {
|
| 92 |
+
# 'model': 'multilingual',
|
| 93 |
+
# 'repo': 'convaiinnovations/laya/multilingual',
|
| 94 |
+
# 'reason': 'non-Latin script (devanagari, 100% of letters); the English checkpoint cannot read it'
|
| 95 |
+
# }
|
| 96 |
+
```
|
| 97 |
+
|
| 98 |
+
### Why Route: The Evidence
|
| 99 |
+
|
| 100 |
+
On a shared benchmark (17,416 questions, one T4 GPU, identical questions per model):
|
| 101 |
+
|
| 102 |
+
| Benchmark / Task | English (`laya`) | Multilingual (`laya-multilingual`) | `Router` (Routed) |
|
| 103 |
+
|---|---|---|---|
|
| 104 |
+
| MASSIVE intent, English | **0.783** | 0.657 | **0.783** |
|
| 105 |
+
| MASSIVE intent, 13 other languages | 0.306 | **0.451** | **0.451** |
|
| 106 |
+
| XNLI, English | **0.860** | 0.843 | **0.860** |
|
| 107 |
+
| XNLI, 14 other languages | 0.521 | **0.731** | **0.731** |
|
| 108 |
+
| Languages usable (>3x random) | 23 / 51 | 45 / 51 | **45 / 51** |
|
| 109 |
+
| Latency, 1 question (T4 GPU) | 39.5 ms | **32.8 ms** | **32.8 ms** |
|
| 110 |
+
| Latency, 10 questions batched | 158.6 ms | **72.3 ms** | **72.3 ms** |
|
| 111 |
+
|
| 112 |
+
The English checkpoint collapses on non-Latin scripts (Khmer scores **0.000 accuracy at 0.952 confidence**). Because the model stays confident while being wrong, confidence gating cannot save you. `Router` detects the script in <0.5 ms pure Python before the forward pass.
|
| 113 |
+
|
| 114 |
+
### Production Preload & Memory
|
| 115 |
+
|
| 116 |
+
A cold checkpoint build costs seconds; language detection costs microseconds. At the default `max_loaded=1`, traffic that alternates languages rebuilds a model on *every* request (measured at a 7.4 s median reload on CPU and 10.3 s on T4).
|
| 117 |
+
|
| 118 |
+
For a server or a demo, preload:
|
| 119 |
+
|
| 120 |
+
```python
|
| 121 |
+
# Every checkpoint resident in memory; language flips cost detection only (<1 ms)
|
| 122 |
+
router = Router(preload=True)
|
| 123 |
+
router = Router(preload=True, device="cuda")
|
| 124 |
+
|
| 125 |
+
# Or preload only the specific checkpoints you serve:
|
| 126 |
+
router.preload(["english", "multilingual"])
|
| 127 |
+
|
| 128 |
+
# If your app already built an agent, attach it to avoid duplicate VRAM:
|
| 129 |
+
router.attach("english", existing_agent)
|
| 130 |
+
|
| 131 |
+
# Manage resident memory (default keeps 1 hot, LRU eviction)
|
| 132 |
+
router = Router(max_loaded=2) # keep two hot
|
| 133 |
+
router.unload() # free memory
|
| 134 |
+
```
|
| 135 |
+
|
| 136 |
+
| Deployment Mode | Per-Request Latency | Model Reloads |
|
| 137 |
+
|---|---|---|
|
| 138 |
+
| `Router()` (lazy, `max_loaded=1`) | 7 to 10 s on every language switch | 1 per switch |
|
| 139 |
+
| `Router(preload=True)` | **32.8 ms (GPU) / 193–464 ms (CPU)** | **none** |
|
| 140 |
+
|
| 141 |
+
---
|
| 142 |
+
|
| 143 |
+
## Single-Model Mode (Direct SDK)
|
| 144 |
+
|
| 145 |
+
If you only need a single checkpoint for a dedicated pipeline:
|
| 146 |
+
|
| 147 |
+
```python
|
| 148 |
+
import laya
|
| 149 |
+
|
| 150 |
+
# 1. Load from the repo root or subfolders (downloads only the requested weights)
|
| 151 |
+
agent = laya.load("convaiinnovations/laya") # English root (~808 MB)
|
| 152 |
+
agent_ml = laya.load("convaiinnovations/laya", subfolder="multilingual") # 100+ languages (~647 MB)
|
| 153 |
+
agent_td = laya.load("convaiinnovations/laya", subfolder="typed-decisions")
|
| 154 |
+
|
| 155 |
+
# 2. Run all questions in ONE single forward pass (~35 ms on GPU)
|
| 156 |
+
result = agent.predict(state, questions)
|
| 157 |
+
answers = result["answers"]
|
| 158 |
+
|
| 159 |
+
print("Department :", answers["department"]["choice"]) # -> billing (confidence: 0.94)
|
| 160 |
+
print("Urgency :", answers["urgency"]["score"]) # -> 1.84 / 2.0
|
| 161 |
+
print("Churn Risk :", answers["churn_risk"]["noul"]) # -> 0.892 (89.2% probability)
|
| 162 |
+
```
|
| 163 |
+
|
| 164 |
+
> **If `laya.load()` hangs:** `transformers` probes for TensorFlow at import, and when TF is
|
| 165 |
+
> installed its abseil runtime can deadlock model construction. Run with `USE_TF=0`.
|
| 166 |
+
|
| 167 |
+
---
|
| 168 |
+
|
| 169 |
+
## Architecture
|
| 170 |
+
|
| 171 |
+
- **Backbone:** ModernBERT-large (395M, bidirectional, fully fine-tuned) + a decision head trained from scratch: 2 transformer layers, an option-marker scorer, and an act/escalate head. 421M total. (Multilingual uses mmBERT-base, 22 layers, 256k vocab, 322M total).
|
| 172 |
+
- **Option markers:** Every option is scored at its own `[MASK]` token, then softmaxed over that question's options. The answer space is defined at request time, so new schemas need no retraining.
|
| 173 |
+
- **Budget:** 512 tokens per question for English (`head_max_len = 192`); 1024 tokens for multilingual (`head_max_len = 256`).
|
| 174 |
+
- **Batching:** Every question in a call is answered in one single forward pass.
|
| 175 |
+
|
| 176 |
+
---
|
| 177 |
+
|
| 178 |
+
## Training
|
| 179 |
+
|
| 180 |
+
**RLCD (Reinforcement Learning for Calibrated Decisions).** The policy reports a distribution; exploration adds zero-mean Gaussian noise to the logits; the reward is a strictly proper scoring rule (log + spherical, plus ranked probability score for ordinal questions). Expected reward is maximised only by reporting honest probabilities. Updates are REINFORCE with a group-mean baseline (GRPO-style). Multi-turn conversations use TD(λ=1.0) over prefix slices.
|
| 181 |
+
|
| 182 |
+
---
|
| 183 |
+
|
| 184 |
+
## Benchmarks
|
| 185 |
+
|
| 186 |
+
Measured on a Tesla T4; every checkpoint answered byte-identical questions in the same run.
|
| 187 |
+
|
| 188 |
+
### Speed
|
| 189 |
+
|
| 190 |
+
| questions per call | `laya` | `laya-multilingual` |
|
| 191 |
+
|---|---|---|
|
| 192 |
+
| 1 | 39.5 ms | **32.8 ms** |
|
| 193 |
+
| 5 | 84.5 ms | **40.1 ms** |
|
| 194 |
+
| 10 | 158.6 ms (15.9 ms/q) | **72.3 ms (7.2 ms/q)** |
|
| 195 |
+
| 50 | 771 ms | **337 ms (6.8 ms/q)** |
|
| 196 |
+
|
| 197 |
+
103–332 questions/sec batched on a single T4. For reference, TypeSafe Jev has been independently measured at 236–276 ms p50 ([AbdelStark](https://github.com/AbdelStark/jev-benchmarks), [nibzard](https://github.com/nibzard/decision-model-benchmark)), so Laya answers a single question roughly **6–8× faster**.
|
| 198 |
+
|
| 199 |
+
### Laya (with routing) vs TypeSafe Jev
|
| 200 |
+
|
| 201 |
+
Every Laya figure is what `Router().predict(...)` returns — the checkpoint the router selects for that input. Jev figures are **third-party published, never measured here** (no TypeSafe API access); sample sizes and prompts differ.
|
| 202 |
+
|
| 203 |
+
| Benchmark / Metric | TypeSafe Jev 1.13.0 | Laya (routed) | Comparison |
|
| 204 |
+
|---|---|---|---|
|
| 205 |
+
| typed-decisions, 2,000 decisions | 0.727 | **0.766** | +0.039 (beats 0.735 teacher ceiling) |
|
| 206 |
+
| AG News, 4 labels | 0.910 | **0.950** | +0.040 |
|
| 207 |
+
| DAIR Emotion, 6 labels | 0.480 | **0.595** | +0.115 |
|
| 208 |
+
| Banking77 (72 vs 77 labels) | **0.870** | 0.425 | Jev leads on >20 options |
|
| 209 |
+
| ECE *(lower better)* | 0.246 | **0.081** | 3× better (post-temperature) |
|
| 210 |
+
| p50 latency, 1 question | 236–276 ms | **32.8 ms** | 7.8× faster |
|
| 211 |
+
| Languages usable (>3x random) | *no published benchmark* | **45 of 51** | Global language coverage |
|
| 212 |
+
| Weights | closed API | **Apache 2.0** | Open weights, on-premise capable |
|
| 213 |
+
| Cost | $0.042 / 1M tokens | **$0 self-hosted** | 100% free |
|
| 214 |
+
|
| 215 |
+
On DAIR Emotion, Jev assigned zero probability to the true label on 16% of examples.
|
| 216 |
+
|
| 217 |
+
#### Where Jev leads
|
| 218 |
+
|
| 219 |
+
* **High-cardinality label spaces (>20 options at default settings):** On Banking77, Jev scores 0.870 (on 72 labels) while Laya scores 0.425 (on 77 labels at default 256-token head budget). Options share a fixed `head_max_len` budget (192 tokens on English, 256 on multilingual), so 77 options receive only ~3 to 4 tokens per label, causing text to become indistinguishable. Jev supports up to 255 options out-of-the-box. While `laya-multilingual` supports 1,024 context (and up to 8,192 in the encoder) and you can raise `agent.cfg["head_max_len"] = 512` at runtime, Jev is currently better suited for 50+ options in a single prompt without tuning.
|
| 220 |
+
* **Soft distribution matching:** On typed-decisions, while Laya achieves higher argmax accuracy (0.766 vs 0.727), Jev achieves higher soft accuracy (0.580 vs 0.471) against the teacher's full probability distributions.
|
| 221 |
+
* **Out-of-the-box raw calibration:** Before temperature scaling, the base checkpoint has higher raw ECE (0.213 vs 0.144). Laya achieves its 0.081 ECE after domain temperature fitting.
|
| 222 |
+
|
| 223 |
+
Full report: [BENCHMARKS.md](https://github.com/NandhaKishorM/laya/blob/main/BENCHMARKS.md).
|
| 224 |
+
|
| 225 |
+
### typed-decisions, measured on all three checkpoints
|
| 226 |
+
|
| 227 |
+
400 cases, 2,000 decisions, four workflows — measured here.
|
| 228 |
+
|
| 229 |
+
| model | accuracy | soft acc | Brier | ECE | score MAE |
|
| 230 |
+
|---|---|---|---|---|---|
|
| 231 |
+
| **`laya-typed-decisions`** | **0.766** | 0.471 | **0.062** | 0.213 | **0.242** |
|
| 232 |
+
| `laya` | 0.362 | 0.332 | 0.316 | 0.175 | 0.694 |
|
| 233 |
+
| `laya-multilingual` | 0.342 | 0.326 | 0.439 | 0.285 | 0.687 |
|
| 234 |
+
| *Jev 1.13.0 (published)* | *0.727* | *0.580* | *0.148* | *0.144* | *0.391* |
|
| 235 |
+
| *teacher self-agreement ceiling* | *0.735* | | | | |
|
| 236 |
+
| *per-question majority class* | *0.461* | | | | |
|
| 237 |
+
|
| 238 |
+
The fine-tuned checkpoint clears the teacher ceiling and wins all four workflows: invoice processing 0.804, security incidents 0.766, customer service 0.764, agent-trace observability 0.730. By primitive: `noul` 0.857, `choice` 0.733, `score` 0.723.
|
| 239 |
+
|
| 240 |
+
The base checkpoints sit below the majority-class baseline here — the capability on this benchmark comes from fine-tuning, which is what the [fine-tuning notebook](https://github.com/NandhaKishorM/laya/blob/main/notebooks/laya_finetune_typed_decisions_2xT4_kaggle.ipynb) is for.
|
| 241 |
+
|
| 242 |
+
---
|
| 243 |
+
|
| 244 |
+
## Honest Limits
|
| 245 |
+
|
| 246 |
+
- **Base checkpoints are near chance on typed-decisions zero-shot** — 0.362 here and 0.352 for multilingual, against a 0.318 random and 0.461 majority-class baseline. The 0.766 belongs to the checkpoint fine-tuned on that benchmark's own training split. Laya is a fast base to specialise, not a zero-shot decision engine.
|
| 247 |
+
- **High-cardinality choice questions and token budgets:** Sequences split into an option prompt budget (`head_max_len`) and the remaining document/state budget (`max_len - head_max_len`):
|
| 248 |
+
* `laya` (English) defaults to 512 context (`head_max_len = 192`, ~320 tokens for state).
|
| 249 |
+
* `laya-multilingual` and `laya-typed-decisions` default to 1,024 context (`head_max_len = 256`, ~768 tokens for state; mmBERT-base encoder supports up to 8,192 with RoPE).
|
| 250 |
+
At default settings, a 77-option question like Banking77 allocates only `(256 - 16) // 77` ≈ 3–4 tokens per label, causing accuracy to fall off sharply (0.425 vs Jev's 0.870). If evaluating 50+ options in a single question:
|
| 251 |
+
1. Raise `agent.cfg["head_max_len"] = 512` and `agent.cfg["max_len"] = 1024` (or up to 2048 / 4096 / 8192) so every option has enough tokens to remain distinct.
|
| 252 |
+
2. Or split large option sets into a two-step coarse-to-fine hierarchical choice.
|
| 253 |
+
- **Ordinal `score` questions are the weakest primitive** (SST-5 0.372).
|
| 254 |
+
- **Ships over-confident:** Refitting one temperature per (question type, option count) moves mean ECE **0.466 → 0.081** (`laya`) and **0.314 → 0.106** (`laya-multilingual`). Do this on your own data before trusting the probabilities.
|
| 255 |
+
- **English only on root:** Use `laya-multilingual` for anything outside English.
|
| 256 |
+
|
| 257 |
+
---
|
| 258 |
+
|
| 259 |
+
## Links
|
| 260 |
+
|
| 261 |
+
- **GitHub:** https://github.com/NandhaKishorM/laya
|
| 262 |
+
- **PyPI:** https://pypi.org/project/laya/
|
| 263 |
+
- **Live Demo:** https://huggingface.co/spaces/convaiinnovations/laya-demo
|
| 264 |
+
- **Write-up:** [Read on Dev.to](https://dev.to/nandakishor_m_6cc0adfde9f/i-built-non-autoregressive-decision-models-a-year-ago-then-a-frontier-lab-called-it-a-18me)
|
| 265 |
+
|
| 266 |
+
Apache 2.0 · Convai Innovations
|
assets/laya_benchmark.png
ADDED
|
Git LFS Details
|
assets/laya_benchmark_common.png
ADDED
|
Git LFS Details
|
assets/laya_vs_jev.png
ADDED
|
Git LFS Details
|
assets/laya_vs_jev_full.png
ADDED
|
Git LFS Details
|
assets/logo-lockup-dark.png
ADDED
|
assets/logo-lockup-dark.svg
ADDED
|
|
assets/logo-lockup.png
ADDED
|
assets/logo-lockup.svg
ADDED
|
|
assets/logo-mark-ink.png
ADDED
|
assets/logo-mark-ink.svg
ADDED
|
|
assets/logo-mark-mono.svg
ADDED
|
|
assets/logo-mark.png
ADDED
|
assets/logo-mark.svg
ADDED
|
|
email_utils.py
ADDED
|
@@ -0,0 +1,71 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Email helpers for RL Agent: clean raw emails into a compact state and a ready-made set of email questions.
|
| 2 |
+
|
| 3 |
+
Jev-style models lose accuracy on long, noisy state, and RL Agent reads at most max_len (512) tokens,
|
| 4 |
+
so strip quoted replies, signatures and disclaimers in code before asking questions.
|
| 5 |
+
"""
|
| 6 |
+
import re
|
| 7 |
+
|
| 8 |
+
_QUOTE_HEADERS = [
|
| 9 |
+
re.compile(r"^\s*On .{0,300}wrote:\s*$", re.I),
|
| 10 |
+
re.compile(r"^\s*-{2,}\s*(Original|Forwarded) Message\s*-{2,}", re.I),
|
| 11 |
+
re.compile(r"^\s*_{8,}\s*$"),
|
| 12 |
+
re.compile(r"^\s*From:\s.+$", re.I),
|
| 13 |
+
]
|
| 14 |
+
_SIGNATURE_MARKERS = [
|
| 15 |
+
re.compile(r"^\s*--\s*$"),
|
| 16 |
+
re.compile(r"^\s*(best|kind|warm|many thanks|thanks|thank you|regards|cheers|sincerely)[\w ,!.]*$", re.I),
|
| 17 |
+
re.compile(r"^\s*sent from my (iphone|android|mobile|ipad)", re.I),
|
| 18 |
+
]
|
| 19 |
+
_DISCLAIMER = re.compile(r"(confidential|intended (solely )?for the (use of the )?(named )?(addressee|recipient)|"
|
| 20 |
+
r"if you (have )?received this (e-?mail|message) in error)", re.I)
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
def clean_email_body(body, max_chars=3000):
|
| 24 |
+
"""Remove quoted history, signature and legal disclaimer; collapse whitespace; truncate."""
|
| 25 |
+
text = (body or "").replace("\r\n", "\n").replace("\r", "\n").replace("\\n", "\n")
|
| 26 |
+
lines = []
|
| 27 |
+
for line in text.split("\n"):
|
| 28 |
+
if any(p.match(line) for p in _QUOTE_HEADERS) and lines:
|
| 29 |
+
break # everything below is the previous thread
|
| 30 |
+
if line.lstrip().startswith(">"):
|
| 31 |
+
continue
|
| 32 |
+
lines.append(line.rstrip())
|
| 33 |
+
# a sign-off only counts near the end (last 40%, or last 8 lines of a short email) and must be a short line
|
| 34 |
+
cut = len(lines)
|
| 35 |
+
for i in range(max(1, min(int(len(lines) * 0.6), len(lines) - 8)), len(lines)):
|
| 36 |
+
if len(lines[i].strip()) <= 40 and any(p.match(lines[i]) for p in _SIGNATURE_MARKERS):
|
| 37 |
+
cut = i
|
| 38 |
+
break
|
| 39 |
+
lines = lines[:cut]
|
| 40 |
+
paragraphs = [p for p in re.split(r"\n\s*\n", "\n".join(lines)) if not _DISCLAIMER.search(p)]
|
| 41 |
+
text = re.sub(r"[ \t]+", " ", "\n\n".join(p.strip() for p in paragraphs if p.strip()))
|
| 42 |
+
return text[:max_chars]
|
| 43 |
+
|
| 44 |
+
|
| 45 |
+
def email_state(subject, body, sender=None, clean=True, **extra):
|
| 46 |
+
"""Build the state dict the email questions refer to (`subject`, `body`, optional `from`)."""
|
| 47 |
+
state = {"subject": (subject or "").strip(), "body": clean_email_body(body) if clean else (body or "")}
|
| 48 |
+
if sender:
|
| 49 |
+
state["from"] = sender
|
| 50 |
+
state.update({k: v for k, v in extra.items() if v is not None})
|
| 51 |
+
return state
|
| 52 |
+
|
| 53 |
+
|
| 54 |
+
def email_questions(categories=None):
|
| 55 |
+
"""A default fan-out of email questions. `categories` = {key: description} for your own routing labels."""
|
| 56 |
+
categories = categories or {
|
| 57 |
+
"billing": "invoices, payments, refunds", "technical": "bugs, outages, integrations",
|
| 58 |
+
"sales": "pricing, demos, new purchases", "account": "login, access, profile changes",
|
| 59 |
+
"hr": "hiring, leave, payroll", "other": "none of the above",
|
| 60 |
+
}
|
| 61 |
+
return {
|
| 62 |
+
"category": {"type": "choice", "instructions": "Which team should handle the email in `body`?", "criteria": categories},
|
| 63 |
+
"is_spam": {"type": "noul", "instructions": "Is this email unsolicited spam or bulk marketing?"},
|
| 64 |
+
"is_phishing": {"type": "noul", "instructions": "Is this email a phishing or scam attempt to steal money, credentials, or personal data?",
|
| 65 |
+
"criteria": {"true": "phishing, scam, or fraud", "false": "a legitimate email"}},
|
| 66 |
+
"urgency": {"type": "score", "instructions": "How urgent is the issue described in `body`?",
|
| 67 |
+
"criteria": ["no time pressure", "needs attention soon", "blocking issue or hard deadline"]},
|
| 68 |
+
"needs_reply": {"type": "noul", "instructions": "Does the sender expect a reply?"},
|
| 69 |
+
"sentiment": {"type": "score", "instructions": "What is the sender's tone in `body`?",
|
| 70 |
+
"criteria": ["angry or very negative", "negative", "neutral", "positive"]},
|
| 71 |
+
}
|
encoder/config.json
ADDED
|
@@ -0,0 +1,84 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"architectures": [
|
| 3 |
+
"ModernBertForMaskedLM"
|
| 4 |
+
],
|
| 5 |
+
"attention_bias": false,
|
| 6 |
+
"attention_dropout": 0.0,
|
| 7 |
+
"bos_token_id": 50281,
|
| 8 |
+
"classifier_activation": "gelu",
|
| 9 |
+
"classifier_bias": false,
|
| 10 |
+
"classifier_dropout": 0.0,
|
| 11 |
+
"classifier_pooling": "mean",
|
| 12 |
+
"cls_token_id": 50281,
|
| 13 |
+
"decoder_bias": true,
|
| 14 |
+
"deterministic_flash_attn": false,
|
| 15 |
+
"dtype": "float32",
|
| 16 |
+
"embedding_dropout": 0.0,
|
| 17 |
+
"eos_token_id": 50282,
|
| 18 |
+
"global_attn_every_n_layers": 3,
|
| 19 |
+
"gradient_checkpointing": false,
|
| 20 |
+
"hidden_activation": "gelu",
|
| 21 |
+
"hidden_size": 1024,
|
| 22 |
+
"initializer_cutoff_factor": 2.0,
|
| 23 |
+
"initializer_range": 0.02,
|
| 24 |
+
"intermediate_size": 2624,
|
| 25 |
+
"layer_norm_eps": 1e-05,
|
| 26 |
+
"layer_types": [
|
| 27 |
+
"full_attention",
|
| 28 |
+
"sliding_attention",
|
| 29 |
+
"sliding_attention",
|
| 30 |
+
"full_attention",
|
| 31 |
+
"sliding_attention",
|
| 32 |
+
"sliding_attention",
|
| 33 |
+
"full_attention",
|
| 34 |
+
"sliding_attention",
|
| 35 |
+
"sliding_attention",
|
| 36 |
+
"full_attention",
|
| 37 |
+
"sliding_attention",
|
| 38 |
+
"sliding_attention",
|
| 39 |
+
"full_attention",
|
| 40 |
+
"sliding_attention",
|
| 41 |
+
"sliding_attention",
|
| 42 |
+
"full_attention",
|
| 43 |
+
"sliding_attention",
|
| 44 |
+
"sliding_attention",
|
| 45 |
+
"full_attention",
|
| 46 |
+
"sliding_attention",
|
| 47 |
+
"sliding_attention",
|
| 48 |
+
"full_attention",
|
| 49 |
+
"sliding_attention",
|
| 50 |
+
"sliding_attention",
|
| 51 |
+
"full_attention",
|
| 52 |
+
"sliding_attention",
|
| 53 |
+
"sliding_attention",
|
| 54 |
+
"full_attention"
|
| 55 |
+
],
|
| 56 |
+
"local_attention": 128,
|
| 57 |
+
"max_position_embeddings": 8192,
|
| 58 |
+
"mlp_bias": false,
|
| 59 |
+
"mlp_dropout": 0.0,
|
| 60 |
+
"model_type": "modernbert",
|
| 61 |
+
"norm_bias": false,
|
| 62 |
+
"norm_eps": 1e-05,
|
| 63 |
+
"num_attention_heads": 16,
|
| 64 |
+
"num_hidden_layers": 28,
|
| 65 |
+
"pad_token_id": 50283,
|
| 66 |
+
"position_embedding_type": "absolute",
|
| 67 |
+
"repad_logits_with_grad": false,
|
| 68 |
+
"rope_parameters": {
|
| 69 |
+
"full_attention": {
|
| 70 |
+
"rope_theta": 160000.0,
|
| 71 |
+
"rope_type": "default"
|
| 72 |
+
},
|
| 73 |
+
"sliding_attention": {
|
| 74 |
+
"rope_theta": 10000.0,
|
| 75 |
+
"rope_type": "default"
|
| 76 |
+
}
|
| 77 |
+
},
|
| 78 |
+
"sep_token_id": 50282,
|
| 79 |
+
"sparse_pred_ignore_index": -100,
|
| 80 |
+
"sparse_prediction": false,
|
| 81 |
+
"tie_word_embeddings": true,
|
| 82 |
+
"transformers_version": "5.0.0",
|
| 83 |
+
"vocab_size": 50368
|
| 84 |
+
}
|
eval/benchmark_comparison.png
ADDED
|
Git LFS Details
|
eval/reliability_eval_in.png
ADDED
|
eval/reliability_eval_zs.png
ADDED
|
eval/results.json
ADDED
|
@@ -0,0 +1,142 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"model": "rl-agent",
|
| 3 |
+
"questions_evaluated": {
|
| 4 |
+
"eval_in": 23024,
|
| 5 |
+
"eval_zs": 2400
|
| 6 |
+
},
|
| 7 |
+
"by_task_family": {
|
| 8 |
+
"eval_in": {
|
| 9 |
+
"conversation outcomes": {
|
| 10 |
+
"n": 3600,
|
| 11 |
+
"accuracy": 0.4817,
|
| 12 |
+
"ece": 0.0193,
|
| 13 |
+
"nll": 0.6933
|
| 14 |
+
},
|
| 15 |
+
"email triage": {
|
| 16 |
+
"n": 2691,
|
| 17 |
+
"accuracy": 0.7321,
|
| 18 |
+
"ece": 0.0172,
|
| 19 |
+
"nll": 0.5954
|
| 20 |
+
},
|
| 21 |
+
"emotion and tone": {
|
| 22 |
+
"n": 1825,
|
| 23 |
+
"accuracy": 0.9058,
|
| 24 |
+
"ece": 0.0183,
|
| 25 |
+
"nll": 0.2382
|
| 26 |
+
},
|
| 27 |
+
"inference and fact checking": {
|
| 28 |
+
"n": 3022,
|
| 29 |
+
"accuracy": 0.8832,
|
| 30 |
+
"ece": 0.054,
|
| 31 |
+
"nll": 0.3404
|
| 32 |
+
},
|
| 33 |
+
"instruction-following tasks": {
|
| 34 |
+
"n": 600,
|
| 35 |
+
"accuracy": 0.8783,
|
| 36 |
+
"ece": 0.0465,
|
| 37 |
+
"nll": 0.3021
|
| 38 |
+
},
|
| 39 |
+
"intent and routing": {
|
| 40 |
+
"n": 1475,
|
| 41 |
+
"accuracy": 0.9912,
|
| 42 |
+
"ece": 0.0085,
|
| 43 |
+
"nll": 0.1811
|
| 44 |
+
},
|
| 45 |
+
"moderation and safety": {
|
| 46 |
+
"n": 2708,
|
| 47 |
+
"accuracy": 0.9671,
|
| 48 |
+
"ece": 0.0613,
|
| 49 |
+
"nll": 0.1527
|
| 50 |
+
},
|
| 51 |
+
"reading comprehension": {
|
| 52 |
+
"n": 770,
|
| 53 |
+
"accuracy": 0.8468,
|
| 54 |
+
"ece": 0.0833,
|
| 55 |
+
"nll": 0.4086
|
| 56 |
+
},
|
| 57 |
+
"response quality scoring": {
|
| 58 |
+
"n": 3146,
|
| 59 |
+
"accuracy": 0.5814,
|
| 60 |
+
"ece": 0.023,
|
| 61 |
+
"nll": 1.0087
|
| 62 |
+
},
|
| 63 |
+
"robustness checks": {
|
| 64 |
+
"n": 744,
|
| 65 |
+
"accuracy": 0.8508,
|
| 66 |
+
"ece": 0.1085,
|
| 67 |
+
"nll": 1.0576
|
| 68 |
+
},
|
| 69 |
+
"search relevance": {
|
| 70 |
+
"n": 733,
|
| 71 |
+
"accuracy": 0.6276,
|
| 72 |
+
"ece": 0.066,
|
| 73 |
+
"nll": 0.7281
|
| 74 |
+
},
|
| 75 |
+
"sentiment and rating": {
|
| 76 |
+
"n": 961,
|
| 77 |
+
"accuracy": 0.4422,
|
| 78 |
+
"ece": 0.4384,
|
| 79 |
+
"nll": 3.545
|
| 80 |
+
},
|
| 81 |
+
"topic classification": {
|
| 82 |
+
"n": 749,
|
| 83 |
+
"accuracy": 0.9386,
|
| 84 |
+
"ece": 0.0285,
|
| 85 |
+
"nll": 0.1957
|
| 86 |
+
}
|
| 87 |
+
},
|
| 88 |
+
"eval_zs": {
|
| 89 |
+
"emotion and tone": {
|
| 90 |
+
"n": 600,
|
| 91 |
+
"accuracy": 0.5833,
|
| 92 |
+
"ece": 0.3178,
|
| 93 |
+
"nll": 1.9761
|
| 94 |
+
},
|
| 95 |
+
"instruction-following tasks": {
|
| 96 |
+
"n": 600,
|
| 97 |
+
"accuracy": 0.8633,
|
| 98 |
+
"ece": 0.0455,
|
| 99 |
+
"nll": 0.3187
|
| 100 |
+
},
|
| 101 |
+
"moderation and safety": {
|
| 102 |
+
"n": 600,
|
| 103 |
+
"accuracy": 0.7967,
|
| 104 |
+
"ece": 0.1713,
|
| 105 |
+
"nll": 1.4151
|
| 106 |
+
},
|
| 107 |
+
"sentiment and rating": {
|
| 108 |
+
"n": 600,
|
| 109 |
+
"accuracy": 0.3617,
|
| 110 |
+
"ece": 0.2915,
|
| 111 |
+
"nll": 1.7985
|
| 112 |
+
}
|
| 113 |
+
}
|
| 114 |
+
},
|
| 115 |
+
"calibration_temperature": [
|
| 116 |
+
1.6369030475616455,
|
| 117 |
+
1.2514300346374512,
|
| 118 |
+
1.983399510383606
|
| 119 |
+
],
|
| 120 |
+
"latency_ms": {
|
| 121 |
+
"1_questions": {
|
| 122 |
+
"p50_ms": 38.4,
|
| 123 |
+
"p95_ms": 42.1
|
| 124 |
+
},
|
| 125 |
+
"10_questions": {
|
| 126 |
+
"p50_ms": 156.0,
|
| 127 |
+
"p95_ms": 158.4
|
| 128 |
+
},
|
| 129 |
+
"50_questions": {
|
| 130 |
+
"p50_ms": 721.4,
|
| 131 |
+
"p95_ms": 733.0
|
| 132 |
+
}
|
| 133 |
+
},
|
| 134 |
+
"act_policy": {
|
| 135 |
+
"eval_in": {
|
| 136 |
+
"automation_rate": 1.0,
|
| 137 |
+
"accuracy_when_acting": 0.8032331136738056,
|
| 138 |
+
"accuracy_when_escalating": null,
|
| 139 |
+
"accuracy_all": 0.8032331136738056
|
| 140 |
+
}
|
| 141 |
+
}
|
| 142 |
+
}
|
eval/results.md
ADDED
|
@@ -0,0 +1,53 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# RL Agent evaluation
|
| 2 |
+
|
| 3 |
+
Metrics after calibration. Zero-shot = task families held out of training entirely.
|
| 4 |
+
|
| 5 |
+
## In-task
|
| 6 |
+
|
| 7 |
+
| task family | questions | accuracy | ECE | NLL |
|
| 8 |
+
|---|---|---|---|---|
|
| 9 |
+
| conversation outcomes | 3600 | 0.482 | 0.019 | 0.693 |
|
| 10 |
+
| email triage | 2691 | 0.732 | 0.017 | 0.595 |
|
| 11 |
+
| emotion and tone | 1825 | 0.906 | 0.018 | 0.238 |
|
| 12 |
+
| inference and fact checking | 3022 | 0.883 | 0.054 | 0.340 |
|
| 13 |
+
| instruction-following tasks | 600 | 0.878 | 0.046 | 0.302 |
|
| 14 |
+
| intent and routing | 1475 | 0.991 | 0.009 | 0.181 |
|
| 15 |
+
| moderation and safety | 2708 | 0.967 | 0.061 | 0.153 |
|
| 16 |
+
| reading comprehension | 770 | 0.847 | 0.083 | 0.409 |
|
| 17 |
+
| response quality scoring | 3146 | 0.581 | 0.023 | 1.009 |
|
| 18 |
+
| robustness checks | 744 | 0.851 | 0.108 | 1.058 |
|
| 19 |
+
| search relevance | 733 | 0.628 | 0.066 | 0.728 |
|
| 20 |
+
| sentiment and rating | 961 | 0.442 | 0.438 | 3.545 |
|
| 21 |
+
| topic classification | 749 | 0.939 | 0.029 | 0.196 |
|
| 22 |
+
|
| 23 |
+
Overall: accuracy 0.753, ECE 0.030, Brier 0.308, accuracy at 50% coverage 0.947
|
| 24 |
+
|
| 25 |
+
## Zero-shot
|
| 26 |
+
|
| 27 |
+
| task family | questions | accuracy | ECE | NLL |
|
| 28 |
+
|---|---|---|---|---|
|
| 29 |
+
| emotion and tone | 600 | 0.583 | 0.318 | 1.976 |
|
| 30 |
+
| instruction-following tasks | 600 | 0.863 | 0.045 | 0.319 |
|
| 31 |
+
| moderation and safety | 600 | 0.797 | 0.171 | 1.415 |
|
| 32 |
+
| sentiment and rating | 600 | 0.362 | 0.291 | 1.798 |
|
| 33 |
+
|
| 34 |
+
Overall: accuracy 0.651, ECE 0.204, Brier 0.532, accuracy at 50% coverage 0.818
|
| 35 |
+
|
| 36 |
+
## Latency
|
| 37 |
+
|
| 38 |
+
```
|
| 39 |
+
{
|
| 40 |
+
"1_questions": {
|
| 41 |
+
"p50_ms": 38.4,
|
| 42 |
+
"p95_ms": 42.1
|
| 43 |
+
},
|
| 44 |
+
"10_questions": {
|
| 45 |
+
"p50_ms": 156.0,
|
| 46 |
+
"p95_ms": 158.4
|
| 47 |
+
},
|
| 48 |
+
"50_questions": {
|
| 49 |
+
"p50_ms": 721.4,
|
| 50 |
+
"p95_ms": 733.0
|
| 51 |
+
}
|
| 52 |
+
}
|
| 53 |
+
```
|
model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:891102d372688fc2a094dac56a384bc537b87c63f21f9f3dac0be2b7cbc8d86c
|
| 3 |
+
size 842609210
|
multilingual/encoder/config.json
ADDED
|
@@ -0,0 +1,79 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"architectures": [
|
| 3 |
+
"ModernBertForMaskedLM"
|
| 4 |
+
],
|
| 5 |
+
"attention_bias": false,
|
| 6 |
+
"attention_dropout": 0.0,
|
| 7 |
+
"bos_token_id": 2,
|
| 8 |
+
"classifier_activation": "gelu",
|
| 9 |
+
"classifier_bias": false,
|
| 10 |
+
"classifier_dropout": 0.0,
|
| 11 |
+
"classifier_pooling": "mean",
|
| 12 |
+
"cls_token_id": 1,
|
| 13 |
+
"decoder_bias": true,
|
| 14 |
+
"deterministic_flash_attn": false,
|
| 15 |
+
"dtype": "float32",
|
| 16 |
+
"embedding_dropout": 0.0,
|
| 17 |
+
"eos_token_id": 1,
|
| 18 |
+
"global_attn_every_n_layers": 3,
|
| 19 |
+
"gradient_checkpointing": false,
|
| 20 |
+
"hidden_activation": "gelu",
|
| 21 |
+
"hidden_size": 768,
|
| 22 |
+
"initializer_cutoff_factor": 2.0,
|
| 23 |
+
"initializer_range": 0.02,
|
| 24 |
+
"intermediate_size": 1152,
|
| 25 |
+
"layer_norm_eps": 1e-05,
|
| 26 |
+
"layer_types": [
|
| 27 |
+
"full_attention",
|
| 28 |
+
"sliding_attention",
|
| 29 |
+
"sliding_attention",
|
| 30 |
+
"full_attention",
|
| 31 |
+
"sliding_attention",
|
| 32 |
+
"sliding_attention",
|
| 33 |
+
"full_attention",
|
| 34 |
+
"sliding_attention",
|
| 35 |
+
"sliding_attention",
|
| 36 |
+
"full_attention",
|
| 37 |
+
"sliding_attention",
|
| 38 |
+
"sliding_attention",
|
| 39 |
+
"full_attention",
|
| 40 |
+
"sliding_attention",
|
| 41 |
+
"sliding_attention",
|
| 42 |
+
"full_attention",
|
| 43 |
+
"sliding_attention",
|
| 44 |
+
"sliding_attention",
|
| 45 |
+
"full_attention",
|
| 46 |
+
"sliding_attention",
|
| 47 |
+
"sliding_attention",
|
| 48 |
+
"full_attention"
|
| 49 |
+
],
|
| 50 |
+
"local_attention": 128,
|
| 51 |
+
"mask_token_id": 4,
|
| 52 |
+
"max_position_embeddings": 8192,
|
| 53 |
+
"mlp_bias": false,
|
| 54 |
+
"mlp_dropout": 0.0,
|
| 55 |
+
"model_type": "modernbert",
|
| 56 |
+
"norm_bias": false,
|
| 57 |
+
"norm_eps": 1e-05,
|
| 58 |
+
"num_attention_heads": 12,
|
| 59 |
+
"num_hidden_layers": 22,
|
| 60 |
+
"pad_token_id": 0,
|
| 61 |
+
"position_embedding_type": "sans_pos",
|
| 62 |
+
"repad_logits_with_grad": false,
|
| 63 |
+
"rope_parameters": {
|
| 64 |
+
"full_attention": {
|
| 65 |
+
"rope_theta": 160000,
|
| 66 |
+
"rope_type": "default"
|
| 67 |
+
},
|
| 68 |
+
"sliding_attention": {
|
| 69 |
+
"rope_theta": 160000,
|
| 70 |
+
"rope_type": "default"
|
| 71 |
+
}
|
| 72 |
+
},
|
| 73 |
+
"sep_token_id": 1,
|
| 74 |
+
"sparse_pred_ignore_index": -100,
|
| 75 |
+
"sparse_prediction": false,
|
| 76 |
+
"tie_word_embeddings": true,
|
| 77 |
+
"transformers_version": "5.0.0",
|
| 78 |
+
"vocab_size": 256000
|
| 79 |
+
}
|
multilingual/model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:9d628fd971b700382ac6f65920a86f149777b2e748e0c955fb3b19695aa8f204
|
| 3 |
+
size 643835514
|
multilingual/rl_agent_config.json
ADDED
|
@@ -0,0 +1,26 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"encoder": "jhu-clsp/mmBERT-base",
|
| 3 |
+
"head_layers": 2,
|
| 4 |
+
"max_len": 1024,
|
| 5 |
+
"head_max_len": 256,
|
| 6 |
+
"max_prefixes": 6,
|
| 7 |
+
"act_costs": {
|
| 8 |
+
"escalate": 0.5
|
| 9 |
+
},
|
| 10 |
+
"cost_wrong_act": 3.0,
|
| 11 |
+
"amp_dtype": "bf16",
|
| 12 |
+
"model_name": "rl-agent",
|
| 13 |
+
"temperature": [
|
| 14 |
+
1.0,
|
| 15 |
+
1.0,
|
| 16 |
+
1.0
|
| 17 |
+
],
|
| 18 |
+
"temperature_by_options": {},
|
| 19 |
+
"training": {
|
| 20 |
+
"updates": 15987,
|
| 21 |
+
"epochs_completed": 4,
|
| 22 |
+
"hours": 4.97,
|
| 23 |
+
"world_size": 1,
|
| 24 |
+
"fine_tuned_from_checkpoint": false
|
| 25 |
+
}
|
| 26 |
+
}
|
multilingual/tokenizer/tokenizer.json
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:609d8f4c067cd3950f88594c5a802616cea245823836ef5848ee4fc40aab5b6f
|
| 3 |
+
size 34363188
|
multilingual/tokenizer/tokenizer_config.json
ADDED
|
@@ -0,0 +1,22 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"bos_token": "<bos>",
|
| 3 |
+
"clean_up_tokenization_spaces": false,
|
| 4 |
+
"cls_token": "<bos>",
|
| 5 |
+
"eos_token": "<eos>",
|
| 6 |
+
"extra_special_tokens": {
|
| 7 |
+
"extra_0": "<start_of_turn>",
|
| 8 |
+
"extra_1": "<end_of_turn>"
|
| 9 |
+
},
|
| 10 |
+
"mask_token": "<mask>",
|
| 11 |
+
"model_input_names": [
|
| 12 |
+
"input_ids",
|
| 13 |
+
"attention_mask"
|
| 14 |
+
],
|
| 15 |
+
"model_max_length": 8192,
|
| 16 |
+
"pad_token": "<pad>",
|
| 17 |
+
"padding_side": "right",
|
| 18 |
+
"sep_token": "<eos>",
|
| 19 |
+
"spaces_between_special_tokens": false,
|
| 20 |
+
"tokenizer_class": "PreTrainedTokenizerFast",
|
| 21 |
+
"unk_token": "<unk>"
|
| 22 |
+
}
|
rl_agent_api.py
ADDED
|
@@ -0,0 +1,78 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Jev-compatible inference for a saved RL Agent model: system_one(state, questions) -> typed answers."""
|
| 2 |
+
import json
|
| 3 |
+
import math
|
| 4 |
+
import os
|
| 5 |
+
|
| 6 |
+
import numpy as np
|
| 7 |
+
import torch
|
| 8 |
+
|
| 9 |
+
from rl_common import (QTYPES, amp_dtype, build_model, build_sequence, collate_items, confidence_from_probs,
|
| 10 |
+
render_options, temp_bucket)
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
class RLAgent:
|
| 14 |
+
def __init__(self, model_dir, device=None):
|
| 15 |
+
from safetensors.torch import load_file
|
| 16 |
+
from transformers import AutoTokenizer
|
| 17 |
+
with open(os.path.join(model_dir, "rl_agent_config.json")) as f:
|
| 18 |
+
self.cfg = json.load(f)
|
| 19 |
+
self.device = torch.device(device or ("cuda" if torch.cuda.is_available() else "cpu"))
|
| 20 |
+
self.tok = AutoTokenizer.from_pretrained(os.path.join(model_dir, "tokenizer"))
|
| 21 |
+
self.model = build_model(self.cfg, encoder_dir=os.path.join(model_dir, "encoder"))
|
| 22 |
+
self.model.load_state_dict(load_file(os.path.join(model_dir, "model.safetensors")), strict=True)
|
| 23 |
+
self.model.to(self.device).eval()
|
| 24 |
+
self.model.encoder.config.reference_compile = False # torch.compile is a loss on small batches / few SMs (T4)
|
| 25 |
+
self.temperature = self.cfg.get("temperature", [1.0, 1.0, 1.0])
|
| 26 |
+
self.temperature_by_options = self.cfg.get("temperature_by_options", {})
|
| 27 |
+
self.dtype = amp_dtype(self.cfg.get("amp_dtype", "fp16"))
|
| 28 |
+
if self.device.type == "cuda" and torch.cuda.get_device_capability(self.device)[0] < 8:
|
| 29 |
+
self.dtype = torch.float16 # e.g. a bf16-trained model evaluated on a T4
|
| 30 |
+
|
| 31 |
+
@staticmethod
|
| 32 |
+
def _to_internal(qdef):
|
| 33 |
+
t = qdef["type"]
|
| 34 |
+
crit = qdef.get("criteria")
|
| 35 |
+
if t == "choice" and isinstance(crit, list):
|
| 36 |
+
crit = {c: None for c in crit}
|
| 37 |
+
return {"t": t, "ins": qdef["instructions"] if isinstance(qdef["instructions"], str) else json.dumps(qdef["instructions"]),
|
| 38 |
+
"crit": crit}
|
| 39 |
+
|
| 40 |
+
@torch.no_grad()
|
| 41 |
+
def system_one(self, state, questions):
|
| 42 |
+
"""questions: {id: {"type": "choice"|"score"|"noul", "instructions": ..., "criteria": ...}} (Jev request shape)."""
|
| 43 |
+
ids, items = list(questions.keys()), []
|
| 44 |
+
for qid in ids:
|
| 45 |
+
q = self._to_internal(questions[qid])
|
| 46 |
+
seq, markers = build_sequence(self.tok, state, q, self.cfg["max_len"], self.cfg["head_max_len"])
|
| 47 |
+
if len(markers) != len(render_options(q)):
|
| 48 |
+
raise ValueError("question %r: options do not fit in head_max_len=%d tokens" % (qid, self.cfg["head_max_len"]))
|
| 49 |
+
items.append({"ids": seq, "markers": markers, "qtype": QTYPES[q["t"]], "target": [0.0] * len(markers), "label": -1,
|
| 50 |
+
"episode": 0, "ep_step": 0, "ep_len": 1, "src": "api"})
|
| 51 |
+
b = collate_items([items], self.tok.pad_token_id)
|
| 52 |
+
use_amp = self.device.type == "cuda"
|
| 53 |
+
with torch.autocast(device_type=self.device.type, dtype=self.dtype, enabled=use_amp):
|
| 54 |
+
logits, act = self.model(b["input_ids"].to(self.device), b["attention_mask"].to(self.device),
|
| 55 |
+
b["marker_pos"].to(self.device), b["marker_mask"].to(self.device), b["qtype"].to(self.device))
|
| 56 |
+
logits, act = logits.float().cpu().numpy(), torch.softmax(act.float(), -1).cpu().numpy()
|
| 57 |
+
answers, n_tokens = {}, int(b["attention_mask"].sum())
|
| 58 |
+
for r, qid in enumerate(ids):
|
| 59 |
+
q = self._to_internal(questions[qid])
|
| 60 |
+
k = len(items[r]["markers"])
|
| 61 |
+
qt = QTYPES[q["t"]]
|
| 62 |
+
z = logits[r, :k] / self.temperature_by_options.get(temp_bucket(qt, k), self.temperature[qt])
|
| 63 |
+
p = np.exp(z - z.max())
|
| 64 |
+
p = p / p.sum()
|
| 65 |
+
ext = {"act_probability": float(act[r, 0])}
|
| 66 |
+
if q["t"] == "choice":
|
| 67 |
+
keys = list(q["crit"].keys())
|
| 68 |
+
answers[qid] = {"type": "choice", "choice": keys[int(p.argmax())],
|
| 69 |
+
"probabilities": {kk: round(float(v), 4) for kk, v in zip(keys, p)},
|
| 70 |
+
"confidence": round(confidence_from_probs(p, k), 4), "rl_agent": ext}
|
| 71 |
+
elif q["t"] == "score":
|
| 72 |
+
answers[qid] = {"type": "score", "score": round(float((np.arange(k) * p).sum()), 4),
|
| 73 |
+
"legend": {str(i): c for i, c in enumerate(q["crit"])},
|
| 74 |
+
"probabilities": {str(i): round(float(v), 4) for i, v in enumerate(p)},
|
| 75 |
+
"confidence": round(confidence_from_probs(p, k), 4), "rl_agent": ext}
|
| 76 |
+
else:
|
| 77 |
+
answers[qid] = {"type": "noul", "noul": round(float(p[1]), 4), "rl_agent": ext}
|
| 78 |
+
return {"model": "rl-agent", "answers": answers, "usage": {"input_tokens": n_tokens, "output_tokens": 0}}
|
rl_agent_config.json
ADDED
|
@@ -0,0 +1,33 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"encoder": "answerdotai/ModernBERT-large",
|
| 3 |
+
"head_layers": 2,
|
| 4 |
+
"max_len": 512,
|
| 5 |
+
"head_max_len": 192,
|
| 6 |
+
"max_prefixes": 6,
|
| 7 |
+
"act_costs": {
|
| 8 |
+
"escalate": 0.5
|
| 9 |
+
},
|
| 10 |
+
"cost_wrong_act": 3.0,
|
| 11 |
+
"amp_dtype": "bf16",
|
| 12 |
+
"model_name": "rl-agent",
|
| 13 |
+
"temperature": [
|
| 14 |
+
1.6369030475616455,
|
| 15 |
+
1.2514300346374512,
|
| 16 |
+
1.983399510383606
|
| 17 |
+
],
|
| 18 |
+
"temperature_by_options": {
|
| 19 |
+
"choice:3-5": 1.7601518630981445,
|
| 20 |
+
"choice:6-10": 1.0000158548355103,
|
| 21 |
+
"score:3-5": 1.2514300346374512,
|
| 22 |
+
"noul:2": 1.983399510383606,
|
| 23 |
+
"choice:11+": 0.10058280825614929,
|
| 24 |
+
"choice:2": 1.9063563346862793
|
| 25 |
+
},
|
| 26 |
+
"training": {
|
| 27 |
+
"updates": 7313,
|
| 28 |
+
"epochs_completed": 1,
|
| 29 |
+
"hours": 1.96,
|
| 30 |
+
"world_size": 1,
|
| 31 |
+
"fine_tuned_from_checkpoint": true
|
| 32 |
+
}
|
| 33 |
+
}
|
rl_common.py
ADDED
|
@@ -0,0 +1,408 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""RL Agent shared code: config, Jev-style question rendering, model, proper-scoring rewards, metrics.
|
| 2 |
+
|
| 3 |
+
Kept Python 3.9 compatible so the same file runs on Kaggle and on a laptop smoke test.
|
| 4 |
+
"""
|
| 5 |
+
import json
|
| 6 |
+
import math
|
| 7 |
+
import os
|
| 8 |
+
import random
|
| 9 |
+
from typing import Dict, List, Optional
|
| 10 |
+
|
| 11 |
+
import numpy as np
|
| 12 |
+
import torch
|
| 13 |
+
import torch.nn as nn
|
| 14 |
+
import torch.nn.functional as F
|
| 15 |
+
import torch.utils.checkpoint
|
| 16 |
+
|
| 17 |
+
QTYPES = {"choice": 0, "score": 1, "noul": 2}
|
| 18 |
+
QTYPE_NAMES = {v: k for k, v in QTYPES.items()}
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
# ----------------------------------------------------------------------------- config
|
| 22 |
+
def load_cfg(path: Optional[str] = None) -> Dict:
|
| 23 |
+
path = path or os.environ.get("RL_AGENT_CFG", "rl_agent_config.json")
|
| 24 |
+
with open(path) as f:
|
| 25 |
+
return json.load(f)
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
# ----------------------------------------------------------------------------- rendering
|
| 29 |
+
def serialize_state(state) -> str:
|
| 30 |
+
if isinstance(state, str):
|
| 31 |
+
return state
|
| 32 |
+
return json.dumps(state, ensure_ascii=False)
|
| 33 |
+
|
| 34 |
+
|
| 35 |
+
def render_options(q: Dict) -> List[str]:
|
| 36 |
+
"""Option texts in label-index order. Noul is always [false, true] so p[1] == noul."""
|
| 37 |
+
t, crit = q["t"], q.get("crit")
|
| 38 |
+
if t == "choice":
|
| 39 |
+
return [k if not v else "%s: %s" % (k, v) for k, v in crit.items()]
|
| 40 |
+
if t == "score":
|
| 41 |
+
return ["level %d: %s" % (i, c) for i, c in enumerate(crit)]
|
| 42 |
+
crit = crit or {}
|
| 43 |
+
return ["false: " + (crit.get("false") or "no, the statement does not hold"),
|
| 44 |
+
"true: " + (crit.get("true") or "yes, the statement holds")]
|
| 45 |
+
|
| 46 |
+
|
| 47 |
+
def build_sequence(tok, state, q: Dict, max_len: int, head_max_len: int,
|
| 48 |
+
option_order: Optional[List[int]] = None, truncate_left: bool = False):
|
| 49 |
+
"""[CLS] <type> instructions [SEP] [MASK] opt0 [MASK] opt1 ... [SEP] state [SEP].
|
| 50 |
+
|
| 51 |
+
Returns input_ids and the positions of the per-option [MASK] markers (in the given option order).
|
| 52 |
+
"""
|
| 53 |
+
mask_tok = tok.mask_token
|
| 54 |
+
opts = render_options(q)
|
| 55 |
+
order = option_order if option_order is not None else list(range(len(opts)))
|
| 56 |
+
ins = str(q["ins"]).replace(mask_tok, " ")
|
| 57 |
+
head_ids = tok("%s question: %s" % (q["t"], ins), add_special_tokens=False)["input_ids"]
|
| 58 |
+
opt_ids = []
|
| 59 |
+
for i in order:
|
| 60 |
+
opt_ids.append([tok.mask_token_id] + tok(" " + opts[i].replace(mask_tok, " "), add_special_tokens=False)["input_ids"][:48])
|
| 61 |
+
opt_budget = head_max_len - sum(len(o) for o in opt_ids)
|
| 62 |
+
if opt_budget < 16: # too many / too long options: shrink every option text evenly
|
| 63 |
+
per = max(4, (head_max_len - 16) // max(1, len(opt_ids)))
|
| 64 |
+
opt_ids = [o[:per] for o in opt_ids]
|
| 65 |
+
opt_budget = head_max_len - sum(len(o) for o in opt_ids)
|
| 66 |
+
head_ids = head_ids[:max(8, opt_budget)]
|
| 67 |
+
ids = [tok.cls_token_id] + head_ids + [tok.sep_token_id]
|
| 68 |
+
markers = []
|
| 69 |
+
for o in opt_ids:
|
| 70 |
+
markers.append(len(ids))
|
| 71 |
+
ids.extend(o)
|
| 72 |
+
ids.append(tok.sep_token_id)
|
| 73 |
+
room = max(0, max_len - len(ids) - 1)
|
| 74 |
+
st = tok(serialize_state(state).replace(mask_tok, " "), add_special_tokens=False)["input_ids"]
|
| 75 |
+
st = st[-room:] if truncate_left else st[:room]
|
| 76 |
+
ids = ids + st + [tok.sep_token_id]
|
| 77 |
+
return ids[:max_len], [m for m in markers if m < max_len]
|
| 78 |
+
|
| 79 |
+
|
| 80 |
+
# ----------------------------------------------------------------------------- model
|
| 81 |
+
class DecisionModel(nn.Module):
|
| 82 |
+
"""Pretrained bidirectional encoder (no LLM, no LoRA) + from-scratch decision head.
|
| 83 |
+
|
| 84 |
+
Each option gets a [MASK] marker; the head scores markers -> softmax over the question's options.
|
| 85 |
+
"""
|
| 86 |
+
|
| 87 |
+
def __init__(self, encoder: nn.Module, head_layers: int = 2, n_act: int = 2, dropout: float = 0.1):
|
| 88 |
+
super().__init__()
|
| 89 |
+
self.encoder = encoder
|
| 90 |
+
d = encoder.config.hidden_size
|
| 91 |
+
nhead = max(1, d // 64)
|
| 92 |
+
layer = nn.TransformerEncoderLayer(d, nhead, 4 * d, dropout, batch_first=True, norm_first=True)
|
| 93 |
+
self.head = nn.TransformerEncoder(layer, head_layers, enable_nested_tensor=False) if head_layers > 0 else None
|
| 94 |
+
self.type_emb = nn.Embedding(3, d)
|
| 95 |
+
self.scorer = nn.Sequential(nn.LayerNorm(d), nn.Linear(d, d), nn.GELU(), nn.Linear(d, 1))
|
| 96 |
+
self.act_head = nn.Sequential(nn.Linear(d + 4, 256), nn.GELU(), nn.Linear(256, n_act))
|
| 97 |
+
self.register_buffer("temperature", torch.ones(3)) # per qtype, fitted post-hoc in evaluate.py
|
| 98 |
+
self.head_checkpointing = False
|
| 99 |
+
|
| 100 |
+
def forward(self, input_ids, attention_mask, marker_pos, marker_mask, qtype, detach_encoder: bool = False):
|
| 101 |
+
h = self.encoder(input_ids=input_ids, attention_mask=attention_mask).last_hidden_state
|
| 102 |
+
if detach_encoder:
|
| 103 |
+
h = h.detach()
|
| 104 |
+
h = h + self.type_emb(qtype)[:, None, :]
|
| 105 |
+
if self.head is not None:
|
| 106 |
+
pad = ~attention_mask.bool()
|
| 107 |
+
for layer in self.head.layers:
|
| 108 |
+
if self.head_checkpointing and self.training and torch.is_grad_enabled():
|
| 109 |
+
h = torch.utils.checkpoint.checkpoint(layer, h, None, pad, use_reentrant=False)
|
| 110 |
+
else:
|
| 111 |
+
h = layer(h, src_key_padding_mask=pad)
|
| 112 |
+
idx = marker_pos.clamp(min=0)[:, :, None].expand(-1, -1, h.size(-1))
|
| 113 |
+
m = torch.gather(h, 1, idx)
|
| 114 |
+
logits = self.scorer(m).squeeze(-1).float()
|
| 115 |
+
logits = logits.masked_fill(~marker_mask, -1e4)
|
| 116 |
+
# act head sees the pooled sequence + detached summary of its own answer distribution
|
| 117 |
+
p = torch.softmax(logits.detach(), -1)
|
| 118 |
+
k = marker_mask.sum(-1).clamp(min=2).float()
|
| 119 |
+
ent = -(p * torch.log(p.clamp_min(1e-9))).sum(-1) / torch.log(k)
|
| 120 |
+
top2 = p.topk(2, -1).values
|
| 121 |
+
feats = torch.stack([top2[:, 0], top2[:, 0] - top2[:, 1], ent, k / 255.0], -1)
|
| 122 |
+
pooled = h[:, 0].float()
|
| 123 |
+
act_logits = self.act_head(torch.cat([pooled, feats], -1))
|
| 124 |
+
return logits, act_logits
|
| 125 |
+
|
| 126 |
+
|
| 127 |
+
def build_model(cfg: Dict, encoder_dir: Optional[str] = None) -> DecisionModel:
|
| 128 |
+
from transformers import AutoConfig, AutoModel
|
| 129 |
+
if encoder_dir: # offline: architecture only, weights come from the saved state dict
|
| 130 |
+
ecfg = AutoConfig.from_pretrained(encoder_dir)
|
| 131 |
+
enc = AutoModel.from_config(ecfg, attn_implementation="sdpa")
|
| 132 |
+
else:
|
| 133 |
+
enc = AutoModel.from_pretrained(cfg["encoder"], attn_implementation="sdpa")
|
| 134 |
+
return DecisionModel(enc, cfg["head_layers"], len(cfg["act_costs"]) + 1)
|
| 135 |
+
|
| 136 |
+
|
| 137 |
+
# ----------------------------------------------------------------------------- rewards (strictly proper)
|
| 138 |
+
def proper_reward(q: torch.Tensor, target: torch.Tensor, qtype: torch.Tensor, mask: torch.Tensor,
|
| 139 |
+
w_sph: float = 0.5, w_rps: float = 1.0, log_floor: float = -9.21) -> torch.Tensor:
|
| 140 |
+
"""q: [..., N, K] reported distributions, target: [N, K] (one-hot or soft) -> reward [..., N].
|
| 141 |
+
|
| 142 |
+
log score + spherical score for all types, + ranked probability score for ordinal (score) questions.
|
| 143 |
+
All three are strictly proper, so the only way to maximize reward is to report honest probabilities.
|
| 144 |
+
"""
|
| 145 |
+
q = q * mask
|
| 146 |
+
logq = torch.log(q.clamp_min(1e-12)).clamp_min(log_floor)
|
| 147 |
+
log_score = (target * logq).sum(-1)
|
| 148 |
+
sph = (target * q).sum(-1) / q.norm(dim=-1).clamp_min(1e-9)
|
| 149 |
+
r = log_score + w_sph * sph
|
| 150 |
+
is_score = (qtype == QTYPES["score"]).float()
|
| 151 |
+
if is_score.any():
|
| 152 |
+
k = mask.sum(-1).clamp(min=2).float()
|
| 153 |
+
cdf_q = torch.cumsum(q, -1)
|
| 154 |
+
cdf_t = torch.cumsum(target, -1)
|
| 155 |
+
rps = (((cdf_q - cdf_t) ** 2) * mask).sum(-1) / (k - 1)
|
| 156 |
+
r = r - w_rps * rps * is_score
|
| 157 |
+
return r
|
| 158 |
+
|
| 159 |
+
|
| 160 |
+
# ----------------------------------------------------------------------------- metrics (numpy, no sklearn)
|
| 161 |
+
def ece_score(conf: np.ndarray, correct: np.ndarray, bins: int = 15) -> float:
|
| 162 |
+
if len(conf) == 0:
|
| 163 |
+
return float("nan")
|
| 164 |
+
edges = np.linspace(0, 1, bins + 1)
|
| 165 |
+
e = 0.0
|
| 166 |
+
for lo, hi in zip(edges[:-1], edges[1:]):
|
| 167 |
+
sel = (conf > lo) & (conf <= hi)
|
| 168 |
+
if sel.any():
|
| 169 |
+
e += sel.mean() * abs(conf[sel].mean() - correct[sel].mean())
|
| 170 |
+
return float(e)
|
| 171 |
+
|
| 172 |
+
|
| 173 |
+
def auroc(scores: np.ndarray, labels: np.ndarray) -> float:
|
| 174 |
+
pos, neg = labels == 1, labels == 0
|
| 175 |
+
if pos.sum() == 0 or neg.sum() == 0:
|
| 176 |
+
return float("nan")
|
| 177 |
+
order = np.argsort(scores)
|
| 178 |
+
ranks = np.empty(len(scores))
|
| 179 |
+
ranks[order] = np.arange(1, len(scores) + 1)
|
| 180 |
+
# average ties
|
| 181 |
+
s_sorted = scores[order]
|
| 182 |
+
i = 0
|
| 183 |
+
while i < len(s_sorted):
|
| 184 |
+
j = i
|
| 185 |
+
while j + 1 < len(s_sorted) and s_sorted[j + 1] == s_sorted[i]:
|
| 186 |
+
j += 1
|
| 187 |
+
if j > i:
|
| 188 |
+
ranks[order[i:j + 1]] = (i + j + 2) / 2.0
|
| 189 |
+
i = j + 1
|
| 190 |
+
return float((ranks[pos].sum() - pos.sum() * (pos.sum() + 1) / 2) / (pos.sum() * neg.sum()))
|
| 191 |
+
|
| 192 |
+
|
| 193 |
+
def spearman(a: np.ndarray, b: np.ndarray) -> float:
|
| 194 |
+
if len(a) < 3:
|
| 195 |
+
return float("nan")
|
| 196 |
+
ra = np.argsort(np.argsort(a)).astype(float)
|
| 197 |
+
rb = np.argsort(np.argsort(b)).astype(float)
|
| 198 |
+
if ra.std() == 0 or rb.std() == 0:
|
| 199 |
+
return float("nan")
|
| 200 |
+
return float(np.corrcoef(ra, rb)[0, 1])
|
| 201 |
+
|
| 202 |
+
|
| 203 |
+
def aurc(conf: np.ndarray, correct: np.ndarray) -> float:
|
| 204 |
+
"""Area under the risk-coverage curve (lower is better)."""
|
| 205 |
+
if len(conf) == 0:
|
| 206 |
+
return float("nan")
|
| 207 |
+
order = np.argsort(-conf)
|
| 208 |
+
err = 1 - correct[order]
|
| 209 |
+
return float((np.cumsum(err) / np.arange(1, len(err) + 1)).mean())
|
| 210 |
+
|
| 211 |
+
|
| 212 |
+
def confidence_from_probs(p: np.ndarray, k: int) -> float:
|
| 213 |
+
"""Jev-style confidence: 1 - normalized entropy of the answer distribution."""
|
| 214 |
+
if k < 2:
|
| 215 |
+
return 1.0
|
| 216 |
+
p = p[:k]
|
| 217 |
+
ent = -(p * np.log(np.clip(p, 1e-12, 1))).sum()
|
| 218 |
+
return float(1 - ent / math.log(k))
|
| 219 |
+
|
| 220 |
+
|
| 221 |
+
def seed_all(seed: int):
|
| 222 |
+
random.seed(seed)
|
| 223 |
+
np.random.seed(seed)
|
| 224 |
+
torch.manual_seed(seed)
|
| 225 |
+
|
| 226 |
+
|
| 227 |
+
# ----------------------------------------------------------------------------- record -> model inputs
|
| 228 |
+
def episode_prefix_lengths(n_turns: int, max_prefixes: int) -> List[int]:
|
| 229 |
+
if n_turns <= max_prefixes:
|
| 230 |
+
return list(range(1, n_turns + 1))
|
| 231 |
+
return sorted(set(int(round(x)) for x in np.linspace(1, n_turns, max_prefixes)))
|
| 232 |
+
|
| 233 |
+
|
| 234 |
+
def encode_record(rec: Dict, tok, cfg: Dict, rng: Optional[random.Random], train: bool) -> List[Dict]:
|
| 235 |
+
"""One stored record -> list of model sequences (one per question, or one per conversation prefix)."""
|
| 236 |
+
items = []
|
| 237 |
+
if rec.get("kind") == "episode":
|
| 238 |
+
ep, q = rec["ep"], rec["qs"][0]
|
| 239 |
+
lens = episode_prefix_lengths(len(ep["turns"]), cfg["max_prefixes"])
|
| 240 |
+
for step, t in enumerate(lens):
|
| 241 |
+
state = dict(ep["ctx"], conversation=ep["turns"][:t])
|
| 242 |
+
ids, markers = build_sequence(tok, state, q, cfg["max_len"], cfg["head_max_len"], truncate_left=True)
|
| 243 |
+
if len(markers) != 2:
|
| 244 |
+
continue
|
| 245 |
+
items.append({"ids": ids, "markers": markers, "qtype": QTYPES["noul"], "target": [1.0 - ep["y"], float(ep["y"])],
|
| 246 |
+
"label": int(ep["y"]), "episode": 1, "ep_step": step, "ep_len": len(lens), "src": rec.get("src", ""),
|
| 247 |
+
"prefix_frac": t / float(len(ep["turns"]))})
|
| 248 |
+
return items
|
| 249 |
+
for qi, q in enumerate(rec["qs"]):
|
| 250 |
+
k = len(render_options(q))
|
| 251 |
+
target = list(q["soft"]) if q.get("soft") else [1.0 if i == q["y"] else 0.0 for i in range(k)]
|
| 252 |
+
order = list(range(k))
|
| 253 |
+
if train and rng is not None and q["t"] != "score":
|
| 254 |
+
rng.shuffle(order)
|
| 255 |
+
ids, markers = build_sequence(tok, rec["state"], q, cfg["max_len"], cfg["head_max_len"], option_order=order)
|
| 256 |
+
if len(markers) != k:
|
| 257 |
+
continue # options did not fit; skip rather than train on a truncated answer space
|
| 258 |
+
target = [target[i] for i in order]
|
| 259 |
+
label = order.index(q["y"]) if q.get("y") is not None else -1
|
| 260 |
+
items.append({"ids": ids, "markers": markers, "qtype": QTYPES[q["t"]], "target": target, "label": label,
|
| 261 |
+
"episode": 0, "ep_step": 0, "ep_len": 1, "src": rec.get("src", ""), "q_index": qi, "order": order})
|
| 262 |
+
return items
|
| 263 |
+
|
| 264 |
+
|
| 265 |
+
def collate_items(batch, pad_id: int):
|
| 266 |
+
items = [it for group in batch for it in group]
|
| 267 |
+
if not items:
|
| 268 |
+
return None
|
| 269 |
+
n, L = len(items), max(len(it["ids"]) for it in items)
|
| 270 |
+
kmax = max(len(it["markers"]) for it in items)
|
| 271 |
+
ids = torch.full((n, L), pad_id, dtype=torch.long)
|
| 272 |
+
att = torch.zeros((n, L), dtype=torch.long)
|
| 273 |
+
mpos = torch.zeros((n, kmax), dtype=torch.long)
|
| 274 |
+
mmask = torch.zeros((n, kmax), dtype=torch.bool)
|
| 275 |
+
target = torch.zeros((n, kmax), dtype=torch.float32)
|
| 276 |
+
ep_group = torch.full((n,), -1, dtype=torch.long)
|
| 277 |
+
group_of = {}
|
| 278 |
+
for i, it in enumerate(items):
|
| 279 |
+
ids[i, :len(it["ids"])] = torch.tensor(it["ids"])
|
| 280 |
+
att[i, :len(it["ids"])] = 1
|
| 281 |
+
k = len(it["markers"])
|
| 282 |
+
mpos[i, :k] = torch.tensor(it["markers"])
|
| 283 |
+
mmask[i, :k] = True
|
| 284 |
+
target[i, :k] = torch.tensor(it["target"], dtype=torch.float32)
|
| 285 |
+
# episodes: all prefixes of the same record share a group id (used for TD(lambda) targets)
|
| 286 |
+
for i, it in enumerate(items):
|
| 287 |
+
if it["episode"]:
|
| 288 |
+
ep_group[i] = group_of.setdefault(it.get("rec_uid", -1 - i), len(group_of))
|
| 289 |
+
return {"input_ids": ids, "attention_mask": att, "marker_pos": mpos, "marker_mask": mmask, "target": target,
|
| 290 |
+
"qtype": torch.tensor([it["qtype"] for it in items]), "label": torch.tensor([it["label"] for it in items]),
|
| 291 |
+
"episode": torch.tensor([it["episode"] for it in items], dtype=torch.bool), "ep_group": ep_group,
|
| 292 |
+
"ep_step": torch.tensor([it["ep_step"] for it in items]), "meta": [{k: it[k] for k in it if k not in ("ids", "markers", "target")} for it in items],
|
| 293 |
+
"n_tokens": int(att.sum())}
|
| 294 |
+
|
| 295 |
+
|
| 296 |
+
def pack_groups(groups: List[List[Dict]], max_tokens: int, max_seqs: int) -> List[List[List[Dict]]]:
|
| 297 |
+
"""Split one sampled batch into sub-batches using the *real* tokenized lengths, so padded tokens never exceed
|
| 298 |
+
max_tokens (the index only stores estimates). A record's items stay together (TD targets need all prefixes)."""
|
| 299 |
+
groups = sorted([g for g in groups if g], key=lambda g: max(len(it["ids"]) for it in g))
|
| 300 |
+
subs, cur, cur_max, cur_n = [], [], 0, 0
|
| 301 |
+
for g in groups:
|
| 302 |
+
g_max, g_n = max(len(it["ids"]) for it in g), len(g)
|
| 303 |
+
if g_max * g_n > max_tokens: # one record bigger than the budget (only if max_tokens < max_len * n_items)
|
| 304 |
+
step = max(1, max_tokens // g_max)
|
| 305 |
+
for s in range(0, g_n, step):
|
| 306 |
+
subs.append([g[s:s + step]])
|
| 307 |
+
continue
|
| 308 |
+
new_max, new_n = max(cur_max, g_max), cur_n + g_n
|
| 309 |
+
if cur and (new_max * new_n > max_tokens or new_n > max_seqs):
|
| 310 |
+
subs.append(cur)
|
| 311 |
+
cur, new_max, new_n = [], g_max, g_n
|
| 312 |
+
cur.append(g)
|
| 313 |
+
cur_max, cur_n = new_max, new_n
|
| 314 |
+
if cur:
|
| 315 |
+
subs.append(cur)
|
| 316 |
+
return subs
|
| 317 |
+
|
| 318 |
+
|
| 319 |
+
def td_lambda_targets(p_true: torch.Tensor, batch: Dict, lam: float) -> torch.Tensor:
|
| 320 |
+
"""TD(lambda) soft targets for conversation prefixes: G_last = outcome, G_t = (1-lam) V_{t+1} + lam G_{t+1}."""
|
| 321 |
+
target = batch["target"].clone()
|
| 322 |
+
groups = batch["ep_group"]
|
| 323 |
+
for g in torch.unique(groups[groups >= 0]).tolist():
|
| 324 |
+
idx = (groups == g).nonzero(as_tuple=True)[0]
|
| 325 |
+
idx = idx[torch.argsort(batch["ep_step"][idx])]
|
| 326 |
+
y = batch["target"][idx[-1], 1]
|
| 327 |
+
G = y
|
| 328 |
+
for j in range(len(idx) - 1, -1, -1):
|
| 329 |
+
if j < len(idx) - 1:
|
| 330 |
+
G = (1 - lam) * p_true[idx[j + 1]] + lam * G
|
| 331 |
+
target[idx[j], 0], target[idx[j], 1] = 1 - G, G
|
| 332 |
+
return target
|
| 333 |
+
|
| 334 |
+
|
| 335 |
+
def make_token_batches(lengths: np.ndarray, nseq: np.ndarray, max_tokens: int, max_seqs: int, rng: np.random.RandomState,
|
| 336 |
+
chunk: int = 4096) -> List[List[int]]:
|
| 337 |
+
"""Length-bucketed batches of record indices under a padded-token budget."""
|
| 338 |
+
order = rng.permutation(len(lengths))
|
| 339 |
+
batches = []
|
| 340 |
+
for s in range(0, len(order), chunk):
|
| 341 |
+
part = order[s:s + chunk]
|
| 342 |
+
part = part[np.argsort(lengths[part])]
|
| 343 |
+
cur, cur_max, cur_n = [], 0, 0
|
| 344 |
+
for i in part:
|
| 345 |
+
ln, ns = int(lengths[i]), int(nseq[i])
|
| 346 |
+
new_max, new_n = max(cur_max, ln), cur_n + ns
|
| 347 |
+
if cur and (new_max * new_n > max_tokens or new_n > max_seqs):
|
| 348 |
+
batches.append(cur)
|
| 349 |
+
cur, new_max, new_n = [], ln, ns
|
| 350 |
+
cur.append(int(i))
|
| 351 |
+
cur_max, cur_n = new_max, new_n
|
| 352 |
+
if cur:
|
| 353 |
+
batches.append(cur)
|
| 354 |
+
rng.shuffle(batches)
|
| 355 |
+
return batches
|
| 356 |
+
|
| 357 |
+
|
| 358 |
+
def temp_bucket(qtype: int, k: int) -> str:
|
| 359 |
+
"""Key for per-cardinality temperature fitting: a 2-option noul and a 20-option choice need different scaling."""
|
| 360 |
+
size = "2" if k <= 2 else "3-5" if k <= 5 else "6-10" if k <= 10 else "11+"
|
| 361 |
+
return "%s:%s" % (QTYPE_NAMES[int(qtype)], size)
|
| 362 |
+
|
| 363 |
+
|
| 364 |
+
def amp_dtype(name: Optional[str]) -> torch.dtype:
|
| 365 |
+
"""'bf16' on GPUs that support it (Ampere+, e.g. RTX 6000 Pro); 'fp16' on T4."""
|
| 366 |
+
return torch.bfloat16 if name == "bf16" else torch.float16
|
| 367 |
+
|
| 368 |
+
|
| 369 |
+
@torch.no_grad()
|
| 370 |
+
def predict_items(model, items: List[Dict], pad_id: int = 0, device=None, max_tokens: int = 16384, use_amp: bool = True,
|
| 371 |
+
dtype: torch.dtype = torch.float16, max_seqs: int = 256, progress: str = ""):
|
| 372 |
+
"""Run the model over pre-encoded items; returns list of dicts with probs/logits (uncalibrated) and act probs."""
|
| 373 |
+
import sys
|
| 374 |
+
import time as _time
|
| 375 |
+
model.eval()
|
| 376 |
+
out = []
|
| 377 |
+
t0, done_tok = _time.time(), 0
|
| 378 |
+
order = sorted(range(len(items)), key=lambda i: len(items[i]["ids"]))
|
| 379 |
+
i = 0
|
| 380 |
+
while i < len(order):
|
| 381 |
+
j, L = i, 0
|
| 382 |
+
while j < len(order) and j - i < max_seqs and max(L, len(items[order[j]]["ids"])) * (j - i + 1) <= max_tokens:
|
| 383 |
+
L = max(L, len(items[order[j]]["ids"]))
|
| 384 |
+
j += 1
|
| 385 |
+
j = max(j, i + 1)
|
| 386 |
+
sel = [items[order[t]] for t in range(i, j)]
|
| 387 |
+
b = collate_items([sel], pad_id)
|
| 388 |
+
with torch.autocast(device_type=device.type, dtype=dtype, enabled=use_amp and device.type == "cuda"):
|
| 389 |
+
logits, act = model(b["input_ids"].to(device), b["attention_mask"].to(device), b["marker_pos"].to(device),
|
| 390 |
+
b["marker_mask"].to(device), b["qtype"].to(device))
|
| 391 |
+
logits, act = logits.float().cpu(), torch.softmax(act.float(), -1).cpu()
|
| 392 |
+
done_tok += int(b["attention_mask"].sum())
|
| 393 |
+
if progress and (j % max(1, len(order) // 2000) == 0 or j >= len(order)):
|
| 394 |
+
el = _time.time() - t0
|
| 395 |
+
eta = el * (len(order) - j) / max(1, j)
|
| 396 |
+
sys.stdout.write("\r [%s] %d/%d sequences | %.1fk tok/s | ETA %dm%02ds " %
|
| 397 |
+
(progress, j, len(order), done_tok / max(el, 1e-9) / 1000, int(eta // 60), int(eta % 60)))
|
| 398 |
+
sys.stdout.flush()
|
| 399 |
+
for r, it in enumerate(sel):
|
| 400 |
+
k = len(it["markers"])
|
| 401 |
+
out.append((order[i + r], {"logits": logits[r, :k].detach().numpy(), "act": act[r].detach().numpy()}))
|
| 402 |
+
i = j
|
| 403 |
+
if progress:
|
| 404 |
+
print("\r [%s] %d sequences in %.0fs (%.1fk tok/s)%s" % (progress, len(order), _time.time() - t0,
|
| 405 |
+
done_tok / max(_time.time() - t0, 1e-9) / 1000, " " * 20))
|
| 406 |
+
out.sort(key=lambda x: x[0])
|
| 407 |
+
model.train()
|
| 408 |
+
return [o for _, o in out]
|
tokenizer/tokenizer.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
tokenizer/tokenizer_config.json
ADDED
|
@@ -0,0 +1,14 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"clean_up_tokenization_spaces": true,
|
| 3 |
+
"cls_token": "[CLS]",
|
| 4 |
+
"mask_token": "[MASK]",
|
| 5 |
+
"model_input_names": [
|
| 6 |
+
"input_ids",
|
| 7 |
+
"attention_mask"
|
| 8 |
+
],
|
| 9 |
+
"model_max_length": 8192,
|
| 10 |
+
"pad_token": "[PAD]",
|
| 11 |
+
"sep_token": "[SEP]",
|
| 12 |
+
"tokenizer_class": "PreTrainedTokenizerFast",
|
| 13 |
+
"unk_token": "[UNK]"
|
| 14 |
+
}
|
typed-decisions/encoder/config.json
ADDED
|
@@ -0,0 +1,84 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"architectures": [
|
| 3 |
+
"ModernBertForMaskedLM"
|
| 4 |
+
],
|
| 5 |
+
"attention_bias": false,
|
| 6 |
+
"attention_dropout": 0.0,
|
| 7 |
+
"bos_token_id": 50281,
|
| 8 |
+
"classifier_activation": "gelu",
|
| 9 |
+
"classifier_bias": false,
|
| 10 |
+
"classifier_dropout": 0.0,
|
| 11 |
+
"classifier_pooling": "mean",
|
| 12 |
+
"cls_token_id": 50281,
|
| 13 |
+
"decoder_bias": true,
|
| 14 |
+
"deterministic_flash_attn": false,
|
| 15 |
+
"dtype": "float32",
|
| 16 |
+
"embedding_dropout": 0.0,
|
| 17 |
+
"eos_token_id": 50282,
|
| 18 |
+
"global_attn_every_n_layers": 3,
|
| 19 |
+
"gradient_checkpointing": false,
|
| 20 |
+
"hidden_activation": "gelu",
|
| 21 |
+
"hidden_size": 1024,
|
| 22 |
+
"initializer_cutoff_factor": 2.0,
|
| 23 |
+
"initializer_range": 0.02,
|
| 24 |
+
"intermediate_size": 2624,
|
| 25 |
+
"layer_norm_eps": 1e-05,
|
| 26 |
+
"layer_types": [
|
| 27 |
+
"full_attention",
|
| 28 |
+
"sliding_attention",
|
| 29 |
+
"sliding_attention",
|
| 30 |
+
"full_attention",
|
| 31 |
+
"sliding_attention",
|
| 32 |
+
"sliding_attention",
|
| 33 |
+
"full_attention",
|
| 34 |
+
"sliding_attention",
|
| 35 |
+
"sliding_attention",
|
| 36 |
+
"full_attention",
|
| 37 |
+
"sliding_attention",
|
| 38 |
+
"sliding_attention",
|
| 39 |
+
"full_attention",
|
| 40 |
+
"sliding_attention",
|
| 41 |
+
"sliding_attention",
|
| 42 |
+
"full_attention",
|
| 43 |
+
"sliding_attention",
|
| 44 |
+
"sliding_attention",
|
| 45 |
+
"full_attention",
|
| 46 |
+
"sliding_attention",
|
| 47 |
+
"sliding_attention",
|
| 48 |
+
"full_attention",
|
| 49 |
+
"sliding_attention",
|
| 50 |
+
"sliding_attention",
|
| 51 |
+
"full_attention",
|
| 52 |
+
"sliding_attention",
|
| 53 |
+
"sliding_attention",
|
| 54 |
+
"full_attention"
|
| 55 |
+
],
|
| 56 |
+
"local_attention": 128,
|
| 57 |
+
"max_position_embeddings": 8192,
|
| 58 |
+
"mlp_bias": false,
|
| 59 |
+
"mlp_dropout": 0.0,
|
| 60 |
+
"model_type": "modernbert",
|
| 61 |
+
"norm_bias": false,
|
| 62 |
+
"norm_eps": 1e-05,
|
| 63 |
+
"num_attention_heads": 16,
|
| 64 |
+
"num_hidden_layers": 28,
|
| 65 |
+
"pad_token_id": 50283,
|
| 66 |
+
"position_embedding_type": "absolute",
|
| 67 |
+
"repad_logits_with_grad": false,
|
| 68 |
+
"rope_parameters": {
|
| 69 |
+
"full_attention": {
|
| 70 |
+
"rope_theta": 160000.0,
|
| 71 |
+
"rope_type": "default"
|
| 72 |
+
},
|
| 73 |
+
"sliding_attention": {
|
| 74 |
+
"rope_theta": 10000.0,
|
| 75 |
+
"rope_type": "default"
|
| 76 |
+
}
|
| 77 |
+
},
|
| 78 |
+
"sep_token_id": 50282,
|
| 79 |
+
"sparse_pred_ignore_index": -100,
|
| 80 |
+
"sparse_prediction": false,
|
| 81 |
+
"tie_word_embeddings": true,
|
| 82 |
+
"transformers_version": "5.17.0",
|
| 83 |
+
"vocab_size": 50368
|
| 84 |
+
}
|
typed-decisions/model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:4fa56de72383a9d3efa9cfa78955733c81b9fc8067a587ca4beb82c78107a24e
|
| 3 |
+
size 842609220
|
typed-decisions/rl_agent_config.json
ADDED
|
@@ -0,0 +1,36 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"encoder": "answerdotai/ModernBERT-large",
|
| 3 |
+
"head_layers": 2,
|
| 4 |
+
"max_len": 1024,
|
| 5 |
+
"head_max_len": 256,
|
| 6 |
+
"max_prefixes": 6,
|
| 7 |
+
"act_costs": {
|
| 8 |
+
"escalate": 0.5
|
| 9 |
+
},
|
| 10 |
+
"cost_wrong_act": 3.0,
|
| 11 |
+
"amp_dtype": "bf16",
|
| 12 |
+
"model_name": "laya-typed-decisions",
|
| 13 |
+
"temperature": [
|
| 14 |
+
1.0148024559020996,
|
| 15 |
+
1.0374259948730469,
|
| 16 |
+
1.0575125217437744
|
| 17 |
+
],
|
| 18 |
+
"temperature_by_options": {
|
| 19 |
+
"choice:3-5": 1.7601518630981445,
|
| 20 |
+
"choice:6-10": 1.0000158548355103,
|
| 21 |
+
"score:3-5": 1.2514300346374512,
|
| 22 |
+
"noul:2": 1.983399510383606,
|
| 23 |
+
"choice:11+": 0.10058280825614929,
|
| 24 |
+
"choice:2": 1.9063563346862793
|
| 25 |
+
},
|
| 26 |
+
"training": {
|
| 27 |
+
"updates": 7313,
|
| 28 |
+
"epochs_completed": 1,
|
| 29 |
+
"hours": 1.96,
|
| 30 |
+
"world_size": 1,
|
| 31 |
+
"fine_tuned_from_checkpoint": true
|
| 32 |
+
},
|
| 33 |
+
"gradient_checkpointing": true,
|
| 34 |
+
"max_tokens_per_batch": 4096,
|
| 35 |
+
"fine_tuned": true
|
| 36 |
+
}
|
typed-decisions/tokenizer/tokenizer.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
typed-decisions/tokenizer/tokenizer_config.json
ADDED
|
@@ -0,0 +1,15 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"clean_up_tokenization_spaces": true,
|
| 3 |
+
"cls_token": "[CLS]",
|
| 4 |
+
"local_files_only": false,
|
| 5 |
+
"mask_token": "[MASK]",
|
| 6 |
+
"model_input_names": [
|
| 7 |
+
"input_ids",
|
| 8 |
+
"attention_mask"
|
| 9 |
+
],
|
| 10 |
+
"model_max_length": 8192,
|
| 11 |
+
"pad_token": "[PAD]",
|
| 12 |
+
"sep_token": "[SEP]",
|
| 13 |
+
"tokenizer_class": "PreTrainedTokenizerFast",
|
| 14 |
+
"unk_token": "[UNK]"
|
| 15 |
+
}
|