sekkit
/

sekkit convaiinnovations commited on
Commit
97f34fe
·
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,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

  • 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
+ }