Deepdive404-3 convaiinnovations commited on
Commit
d236df4
·
0 Parent(s):

Duplicate from convaiinnovations/laya

Browse files

Co-authored-by: Convai Innovations <convaiinnovations@users.noreply.huggingface.co>

.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,384 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
+ ## Installation
13
+
14
+ ```bash
15
+ pip install laya
16
+ ```
17
+
18
+ Python 3.10 or newer. Optional extras: `laya[serve]` (HTTP server), `laya[mcp]` (MCP server), `laya[langchain]` (LangChain and LangGraph), `laya[onnx]` (ONNX Runtime), `laya[fast]` (TileLang GPU fast path). Platform-by-platform setup is in the [GitHub README](https://github.com/NandhaKishorM/laya#installation-details).
19
+
20
+ **Long documents.** `laya-multilingual` reads up to 8,192 tokens with `max_len=8192`. Measured accuracy and time by document length ([benchmark script](https://github.com/NandhaKishorM/laya/blob/main/research/scripts/bench_long_context.py)):
21
+
22
+ <p align="center">
23
+ <img src="https://raw.githubusercontent.com/NandhaKishorM/laya/main/assets/long_context_8192.png" alt="laya-multilingual long-document accuracy by document length" width="100%" />
24
+ </p>
25
+
26
+ ## Quickstart
27
+
28
+ > **Long documents: `laya-multilingual` reads up to 8,192 tokens.** It ships with a 1,024-token limit that cuts long documents off, so pass `max_len=8192` for them:
29
+ >
30
+ > ```python
31
+ > result = router.predict(long_document, questions, model="multilingual", max_len=8192)
32
+ > ```
33
+ >
34
+ > In the table above, 16 to 18 of 20 requests were answered correctly with up to about 4,000 tokens of text before them; beyond that results vary (8 to 17 of 20), so check long-document accuracy on your own data. Short inputs give identical answers with `max_len=8192`, and speed follows the input's real length, not the limit: short inputs are unchanged, and a 4,000-token input takes about 1.7 s on an Apple GPU. Name the checkpoint with `model="multilingual"`, since long mostly-English text would otherwise route to the English checkpoint.
35
+
36
+ ```python
37
+ from laya import Router
38
+
39
+ router = Router() # downloads a checkpoint on first use; Router(preload=True) loads all three up front
40
+
41
+ state = "Hi, we were billed twice for March. Please refund the duplicate today or we will cancel our plan."
42
+ questions = {
43
+ "department": {"type": "choice", "instructions": "Which department should handle this?",
44
+ "criteria": {"billing": "invoices, payments, refunds",
45
+ "technical": "bugs, outages, system errors",
46
+ "other": "everything else"}},
47
+ "urgency": {"type": "score", "instructions": "How urgent is this?",
48
+ "criteria": ["not urgent", "soon", "blocking"]},
49
+ "churn_risk": {"type": "noul", "instructions": "Does the user threaten to cancel or leave?"},
50
+ }
51
+
52
+ result = router.predict(state, questions)
53
+ print(result["answers"]["department"]["choice"]) # billing
54
+ print(result["answers"]["churn_risk"]["noul"]) # probability the answer is yes
55
+ print(result["routing"]["model"]) # english
56
+ ```
57
+
58
+ The same call works in any of 100+ languages. The `Router` detects the script and language and sends non-English text to `laya-multilingual`:
59
+
60
+ ```python
61
+ for text in ["मुझसे मार्च में दो बार शुल्क लिया गया, कृपया डुप्लिकेट राशि वापस करें।",
62
+ "La aplicación se cierra cada vez que abro la configuración."]:
63
+ r = router.predict(text, {"department": questions["department"]})
64
+ print(r["routing"]["model"], r["answers"]["department"]["choice"])
65
+ # multilingual billing
66
+ # multilingual technical
67
+ ```
68
+
69
+ ## Fine-tune for better accuracy
70
+
71
+ The shipped checkpoints work zero-shot, but fine-tuning on decisions from your own domain is where accuracy jumps. On the typed-decisions benchmark (2,000 decisions across four workflows), the fine-tuned [`laya-typed-decisions`](https://huggingface.co/convaiinnovations/laya-typed-decisions) checkpoint scores **0.766** accuracy, against **0.362** for the base English checkpoint on the same decisions.
72
+
73
+ **[Fine-tuning notebook](https://github.com/NandhaKishorM/laya/blob/main/notebooks/laya_finetune_typed_decisions_2xT4_kaggle.ipynb)**: runs the whole loop on Kaggle's free 2x T4 GPUs (build the dataset, train, fit calibration temperatures, evaluate, and push the result to the Hub). Details in the [GitHub README](https://github.com/NandhaKishorM/laya#fine-tuning).
74
+
75
+ ## Documentation
76
+
77
+ **[nandhakishorm.github.io/laya](https://nandhakishorm.github.io/laya/)**: guides for [prediction hooks](https://nandhakishorm.github.io/laya/hooks/), [schema-driven decisions](https://nandhakishorm.github.io/laya/structured/), [Docker](https://nandhakishorm.github.io/laya/docker/) and [LangChain and LangGraph](https://nandhakishorm.github.io/laya/langchain/), plus a full [API reference](https://nandhakishorm.github.io/laya/reference/).
78
+
79
+ ## What's new in laya 0.3.20
80
+
81
+ The checkpoints themselves are unchanged. `pip install -U laya` for the latest runtime fixes:
82
+
83
+ * **Long documents with `max_len=8192`** on `laya-multilingual`: a measured table shows accuracy and time by document length.
84
+
85
+ * **Sturdier fast path.** After a CUDA out-of-memory error the fallback to CPU switches the TileLang fast path off first, a single-option `choice` no longer crashes it, and concurrent calls can no longer overwrite each other's CUDA-graph buffers.
86
+ * **Server and runtime.** `laya-serve` drains its inference pool on shutdown and returns 401 for a malformed bearer header, and `ONNXAgent` matches `Agent` on empty question sets and long conversation lists.
87
+
88
+ ---
89
+
90
+ <p align="center">
91
+ <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%" />
92
+ </p>
93
+
94
+ **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:
95
+
96
+ | Checkpoint | Backbone Encoder | Params | Context | Best at |
97
+ |---|---|---|---|---|
98
+ | **`convaiinnovations/laya`** (this repo root) | ModernBERT-large | 421M | 512 | English text, guardrails, email triage |
99
+ | [`convaiinnovations/laya-multilingual`](https://huggingface.co/convaiinnovations/laya-multilingual) | mmBERT-base | 322M | 1024 (up to 8k) | 100+ languages, ~2.2x faster |
100
+ | [`convaiinnovations/laya-typed-decisions`](https://huggingface.co/convaiinnovations/laya-typed-decisions) | ModernBERT-large | 421M | 1024 | the four typed-decisions workflows (0.766 acc) |
101
+
102
+
103
+ ## Quickstart: Route Mode (Recommended)
104
+
105
+ 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.
106
+
107
+ ```bash
108
+ pip install laya
109
+ ```
110
+
111
+ ```python
112
+ import laya
113
+ from laya import Router
114
+
115
+ # Preload checkpoints into memory for instant sub-35ms routing
116
+ router = Router(preload=True)
117
+
118
+ state = {
119
+ "from": "user@acme.com",
120
+ "subject": "Duplicate charge on invoice #4411",
121
+ "body": "Hi, we were billed twice for March. Please refund the duplicate today or we will cancel our plan."
122
+ }
123
+
124
+ questions = {
125
+ "department": {
126
+ "type": "choice",
127
+ "instructions": "Which department should handle this request?",
128
+ "criteria": {
129
+ "billing": "invoices, payments, refunds",
130
+ "technical": "bugs, outages, system errors",
131
+ "sales": "pricing, new contracts",
132
+ "other": "everything else"
133
+ }
134
+ },
135
+ "urgency": {
136
+ "type": "score",
137
+ "instructions": "How urgent is this request?",
138
+ "criteria": ["not urgent", "soon", "critical deadline or blocking issue"]
139
+ },
140
+ "churn_risk": {
141
+ "type": "noul",
142
+ "instructions": "Does the user threaten to cancel or leave?"
143
+ },
144
+ "refund_requested": {
145
+ "type": "noul",
146
+ "instructions": "Does the user explicitly request a refund?"
147
+ }
148
+ }
149
+
150
+ # 1. English state -> automatically routed to ModernBERT-large (39.5 ms)
151
+ res_en = router.predict(state, questions)
152
+ print("Department :", res_en["answers"]["department"]["choice"]) # -> billing (confidence: 0.94)
153
+ print("Routing :", res_en["routing"]["model"]) # -> english
154
+
155
+ # 2. Hindi state -> automatically routed to mmBERT-base (100+ languages, 32.8 ms)
156
+ res_hi = router.predict({"body": "मुझसे दो बार शुल्क लिया गया, कृपया पैसे वापस करें।"}, questions)
157
+ print("Department :", res_hi["answers"]["department"]["choice"]) # -> billing (confidence: 0.86)
158
+ print("Routing :", res_hi["routing"]["model"]) # -> multilingual
159
+
160
+ # 3. Explicit override when you already know the checkpoint
161
+ res_td = router.predict(state, questions, model="typed-decisions")
162
+ ```
163
+
164
+ Every result carries full routing metadata explaining why the choice was made:
165
+
166
+ ```python
167
+ res_hi["routing"]
168
+ # {
169
+ # 'model': 'multilingual',
170
+ # 'repo': 'convaiinnovations/laya/multilingual',
171
+ # 'reason': 'non-Latin script (devanagari, 100% of letters); the English checkpoint cannot read it'
172
+ # }
173
+ ```
174
+
175
+ ### Why Route: The Evidence
176
+
177
+ On a shared benchmark (17,416 questions, one T4 GPU, identical questions per model):
178
+
179
+ | Benchmark / Task | English (`laya`) | Multilingual (`laya-multilingual`) | `Router` (Routed) |
180
+ |---|---|---|---|
181
+ | MASSIVE intent, English | **0.783** | 0.657 | **0.783** |
182
+ | MASSIVE intent, 13 other languages | 0.306 | **0.451** | **0.451** |
183
+ | XNLI, English | **0.860** | 0.843 | **0.860** |
184
+ | XNLI, 14 other languages | 0.521 | **0.731** | **0.731** |
185
+ | Languages usable (>3x random) | 23 / 51 | 45 / 51 | **45 / 51** |
186
+ | Latency, 1 question (T4 GPU) | 39.5 ms | **32.8 ms** | **32.8 ms** |
187
+ | Latency, 10 questions batched | 158.6 ms | **72.3 ms** | **72.3 ms** |
188
+
189
+ 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.
190
+
191
+ ### Supplying your own language detection
192
+
193
+ If you already run a language-identification model, pass its answer instead of relying on the built-in heuristic. `lang_guess` takes a language code or a callable, is checked after an explicit `lang=` and before detection, and a callable that returns `None` falls through to detection:
194
+
195
+ ```python
196
+ router.predict(state, questions, lang_guess="ro") # a code you already know
197
+ router = Router(preload=True, lang_guess=my_lid) # or install one for every request
198
+ ```
199
+
200
+ ### Production Preload & Memory
201
+
202
+ A cold checkpoint build costs seconds; language detection costs microseconds. The lazy default keeps **two** checkpoints resident (`english` and `multilingual`, the only two automatic routing chooses between), so after each language's first load a switch costs detection only. A single-language deployment never builds the second. `max_loaded=1` rebuilds on every switch (measured at a 7.4 s median reload on CPU and 10.3 s on T4).
203
+
204
+ For a server or a demo, preload:
205
+
206
+ ```python
207
+ # Every checkpoint resident in memory; language flips cost detection only (<1 ms)
208
+ router = Router(preload=True)
209
+ router = Router(preload=True, device="cuda")
210
+
211
+ # Or preload only the specific checkpoints you serve:
212
+ router.preload(["english", "multilingual"])
213
+
214
+ # If your app already built an agent, attach it to avoid duplicate VRAM:
215
+ router.attach("english", existing_agent)
216
+
217
+ # Manage resident memory (default keeps two hot: english + multilingual, LRU eviction)
218
+ router = Router(max_loaded=3) # all three hot, e.g. with auto_task_detection
219
+ router = Router(max_loaded=1) # memory-constrained host, reloads on every switch
220
+ router.unload() # free memory
221
+
222
+ with Router() as r: # releases the models when the block ends
223
+ r.predict(state, questions)
224
+ ```
225
+
226
+ | Deployment Mode | Per-Request Latency | Model Reloads |
227
+ |---|---|---|
228
+ | `Router()` (lazy, `max_loaded=2`) | detection only (<1 ms) after each language's first load | 1 the first time a language appears |
229
+ | `Router(max_loaded=1)` | 7 to 10 s on every language switch | 1 per switch |
230
+ | `Router(preload=True)` | **32.8 ms (GPU) / 193–464 ms (CPU)** | **none** |
231
+
232
+ ---
233
+
234
+ ## Single-Model Mode (Direct SDK)
235
+
236
+ If you only need a single checkpoint for a dedicated pipeline:
237
+
238
+ ```python
239
+ import laya
240
+
241
+ # 1. Load from the repo root or subfolders (downloads only the requested weights)
242
+ agent = laya.load("convaiinnovations/laya") # English root (~808 MB)
243
+ agent_ml = laya.load("convaiinnovations/laya", subfolder="multilingual") # 100+ languages (~647 MB)
244
+ agent_td = laya.load("convaiinnovations/laya", subfolder="typed-decisions")
245
+
246
+ # 2. Run all questions in ONE single forward pass (~35 ms on GPU)
247
+ result = agent.predict(state, questions)
248
+ answers = result["answers"]
249
+
250
+ print("Department :", answers["department"]["choice"]) # -> billing (confidence: 0.94)
251
+ print("Urgency :", answers["urgency"]["score"]) # -> 1.84 / 2.0
252
+ print("Churn Risk :", answers["churn_risk"]["noul"]) # -> 0.892 (89.2% probability)
253
+ ```
254
+
255
+ > **If `laya.load()` hangs:** `transformers` probes for TensorFlow at import, and when TF is
256
+ > installed its abseil runtime can deadlock model construction. Run with `USE_TF=0`.
257
+
258
+ ---
259
+
260
+ ## Self-hosting: Jev-compatible HTTP server
261
+
262
+ `laya-serve` exposes the `Router` on the same `POST /v1/systemone` request and response shape as TypeSafe Jev, so existing TypeSafe clients work by changing their base URL:
263
+
264
+ ```bash
265
+ pip install "laya[serve]"
266
+ LAYA_DEVICE=cuda LAYA_PRELOAD=1 laya-serve # 0.0.0.0:8000, preloads the checkpoints
267
+ ```
268
+
269
+ ```bash
270
+ curl -s localhost:8000/v1/systemone -H 'Content-Type: application/json' -d '{
271
+ "state": {"document": "I was charged twice. Please fix this ASAP."},
272
+ "questions": {"billing": {"type": "noul", "instructions": "Is this ticket about billing?"}}
273
+ }'
274
+ ```
275
+
276
+ It accepts every question shape the Jev API does (for example `criteria` as a list), ignores unknown fields, and returns a 422 naming the problem for a malformed question. It binds `0.0.0.0` with no authentication unless `LAYA_API_KEY` is set, in which case it requires `Authorization: Bearer <key>`.
277
+
278
+ ---
279
+
280
+ ## Architecture
281
+
282
+ - **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).
283
+ - **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.
284
+ - **Budget:** 512 tokens per question for English (`head_max_len = 192`); 1024 tokens for multilingual (`head_max_len = 256`).
285
+ - **Batching:** Every question in a call is answered in one single forward pass.
286
+
287
+ ---
288
+
289
+ ## Training
290
+
291
+ **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.
292
+
293
+ ---
294
+
295
+ ## Benchmarks
296
+
297
+ Measured on a Tesla T4; every checkpoint answered byte-identical questions in the same run.
298
+
299
+ ### Speed
300
+
301
+ | questions per call | `laya` | `laya-multilingual` |
302
+ |---|---|---|
303
+ | 1 | 39.5 ms | **32.8 ms** |
304
+ | 5 | 84.5 ms | **40.1 ms** |
305
+ | 10 | 158.6 ms (15.9 ms/q) | **72.3 ms (7.2 ms/q)** |
306
+ | 50 | 771 ms | **337 ms (6.8 ms/q)** |
307
+
308
+ 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**.
309
+
310
+ ### Laya (with routing) vs TypeSafe Jev
311
+
312
+ 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.
313
+
314
+ | Benchmark / Metric | TypeSafe Jev 1.13.0 | Laya (routed) | Comparison |
315
+ |---|---|---|---|
316
+ | typed-decisions, 2,000 decisions | 0.727 | **0.766** | +0.039 (beats 0.735 teacher ceiling) |
317
+ | AG News, 4 labels | 0.910 | **0.950** | +0.040 |
318
+ | DAIR Emotion, 6 labels | 0.480 | **0.595** | +0.115 |
319
+ | Banking77 (72 vs 77 labels) | **0.870** | 0.425 | Jev leads on >20 options |
320
+ | ECE *(lower better)* | 0.246 | **0.081** | 3× better (post-temperature) |
321
+ | p50 latency, 1 question | 236–276 ms | **32.8 ms** | 7.8× faster |
322
+ | Languages usable (>3x random) | *no published benchmark* | **45 of 51** | Global language coverage |
323
+ | Weights | closed API | **Apache 2.0** | Open weights, on-premise capable |
324
+ | Cost | $0.042 / 1M tokens | **$0 self-hosted** | 100% free |
325
+
326
+ On DAIR Emotion, Jev assigned zero probability to the true label on 16% of examples.
327
+
328
+ #### Where Jev leads
329
+
330
+ * **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.
331
+ * **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.
332
+ * **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.
333
+
334
+ Full report: [BENCHMARKS.md](https://github.com/NandhaKishorM/laya/blob/main/BENCHMARKS.md).
335
+
336
+ ### typed-decisions, measured on all three checkpoints
337
+
338
+ 400 cases, 2,000 decisions, four workflows — measured here.
339
+
340
+ | model | accuracy | soft acc | Brier | ECE | score MAE |
341
+ |---|---|---|---|---|---|
342
+ | **`laya-typed-decisions`** | **0.766** | 0.471 | **0.062** | 0.213 | **0.242** |
343
+ | `laya` | 0.362 | 0.332 | 0.316 | 0.175 | 0.694 |
344
+ | `laya-multilingual` | 0.342 | 0.326 | 0.439 | 0.285 | 0.687 |
345
+ | *Jev 1.13.0 (published)* | *0.727* | *0.580* | *0.148* | *0.144* | *0.391* |
346
+ | *teacher self-agreement ceiling* | *0.735* | | | | |
347
+ | *per-question majority class* | *0.461* | | | | |
348
+
349
+ 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.
350
+
351
+ 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.
352
+
353
+ ---
354
+
355
+ ## Honest Limits
356
+
357
+ - **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.
358
+ - **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`):
359
+ * `laya` (English) defaults to 512 context (`head_max_len = 192`, ~320 tokens for state).
360
+ * `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).
361
+ 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:
362
+ 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.
363
+ 2. Or split large option sets into a two-step coarse-to-fine hierarchical choice.
364
+ - **Ordinal `score` questions are the weakest primitive** (SST-5 0.372).
365
+ - **`noul` can follow its option labels instead of the state, most strongly on this English checkpoint.** `noul` renders its two options as `false:` / `true:`, and here that label pair can dominate the answer, returning a confident "no" for clearly positive input ([#156](https://github.com/NandhaKishorM/laya/issues/156)). Check `noul` answers on your own data. If they look stuck, ask the same question as a two-option `choice` with neutral keys and your yes/no wording as the descriptions:
366
+
367
+ ```python
368
+ {"type": "choice", "instructions": "Is this review positive?",
369
+ "criteria": {"A": "yes, the review is positive", "B": "no, the review is negative"}}
370
+ ```
371
+ - **`action.act_probability` carries no usable signal yet** ([#185](https://github.com/NandhaKishorM/laya/issues/185)). It reads 1.0 for almost every input, and its raw logits run against correctness (AUROC 0.30 on 396 labelled decisions). Gate on `confidence` instead, which reaches an AUROC of 0.77 on the same items.
372
+ - **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.
373
+ - **English only on root:** Use `laya-multilingual` for anything outside English.
374
+
375
+ ---
376
+
377
+ ## Links
378
+
379
+ - **GitHub:** https://github.com/NandhaKishorM/laya
380
+ - **PyPI:** https://pypi.org/project/laya/
381
+ - **Live Demo:** https://huggingface.co/spaces/convaiinnovations/laya-demo
382
+ - **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)
383
+
384
+ Apache 2.0 · Convai Innovations
assets/laya_benchmark.png ADDED

Git LFS Details

  • SHA256: e01e49f0d842b4616e44c4c9a0feb92a9ae45efeb533f8f838888715dc89c2b9
  • Pointer size: 131 Bytes
  • Size of remote file: 261 kB
assets/laya_benchmark_common.png ADDED

Git LFS Details

  • SHA256: 183b0b17e8d90e091582e415d9f1da0ef34b40c3db6ccd2c8bbfd2b9232b95a6
  • Pointer size: 131 Bytes
  • Size of remote file: 278 kB
assets/laya_vs_jev.png ADDED

Git LFS Details

  • SHA256: 5c06517ea7f3e5f3f84873ddaa3cb470f101ad0276f921c58886f8fb12fbe0a8
  • Pointer size: 131 Bytes
  • Size of remote file: 216 kB
assets/laya_vs_jev_full.png ADDED

Git LFS Details

  • SHA256: ee47b751d524a0bb65159e026141b4ce70cc83d03b14df8ec68afbc7bcdb95b3
  • Pointer size: 131 Bytes
  • Size of remote file: 487 kB
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

  • SHA256: 78e4b926bfca3950ff5441d1030a1453515973fbe5d7dd70a1c764a81f0d3d38
  • Pointer size: 131 Bytes
  • Size of remote file: 467 kB
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
+ }