Zero-Shot Classification
ONNX
Safetensors
MLX
English
intent-classification
intent-detection
text-classification
chatbot
conversational-ai
customer-support
routing
out-of-scope-detection
int8
cpu
modernbert
ettin
laya
distillation
knowledge-distillation
onnxruntime
apple-silicon
macos
metal
on-device
zero-shot
nlu
intent-router
semantic-router
llm-router
open-intent-detection
out-of-distribution-detection
customer-service
banking
e-commerce
edge
Eval Results (legacy)
Instructions to use vrajnotviraj/laya-intent-router-150m-onnx with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- MLX
How to use vrajnotviraj/laya-intent-router-150m-onnx with MLX:
# Download the model from the Hub pip install huggingface_hub[hf_xet] huggingface-cli download --local-dir laya-intent-router-150m-onnx vrajnotviraj/laya-intent-router-150m-onnx
- Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- LM Studio
- Atomic Chat
Laya Intent Router 150M: int8 ONNX, shortlist embedder, router.py, model card
Browse files- .gitattributes +1 -0
- README.md +182 -0
- laya.onnx +3 -0
- laya.onnx.data +3 -0
- rl_agent_config.json +27 -0
- router.py +203 -0
- shortlist/model.onnx +3 -0
- shortlist/shortlist_config.json +1 -0
- shortlist/tokenizer.json +0 -0
- tokenizer.json +0 -0
.gitattributes
CHANGED
|
@@ -33,3 +33,4 @@ saved_model/**/* 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
|
|
|
|
|
|
| 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 |
+
laya.onnx.data filter=lfs diff=lfs merge=lfs -text
|
README.md
ADDED
|
@@ -0,0 +1,182 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
---
|
| 2 |
+
license: apache-2.0
|
| 3 |
+
language:
|
| 4 |
+
- en
|
| 5 |
+
library_name: onnx
|
| 6 |
+
pipeline_tag: zero-shot-classification
|
| 7 |
+
base_model: jhu-clsp/ettin-encoder-150m
|
| 8 |
+
tags:
|
| 9 |
+
- intent-classification
|
| 10 |
+
- intent-detection
|
| 11 |
+
- zero-shot-classification
|
| 12 |
+
- text-classification
|
| 13 |
+
- chatbot
|
| 14 |
+
- conversational-ai
|
| 15 |
+
- customer-support
|
| 16 |
+
- routing
|
| 17 |
+
- out-of-scope-detection
|
| 18 |
+
- onnx
|
| 19 |
+
- int8
|
| 20 |
+
- cpu
|
| 21 |
+
- modernbert
|
| 22 |
+
- ettin
|
| 23 |
+
- laya
|
| 24 |
+
- distillation
|
| 25 |
+
datasets:
|
| 26 |
+
- clinc_oos
|
| 27 |
+
- PolyAI/banking77
|
| 28 |
+
- FastFit/hwu_64
|
| 29 |
+
- benayas/snips
|
| 30 |
+
- bitext/Bitext-customer-support-llm-chatbot-training-dataset
|
| 31 |
+
- bitext/Bitext-retail-ecommerce-llm-chatbot-training-dataset
|
| 32 |
+
- bitext/Bitext-retail-banking-llm-chatbot-training-dataset
|
| 33 |
+
metrics:
|
| 34 |
+
- accuracy
|
| 35 |
+
model-index:
|
| 36 |
+
- name: laya-intent-router-150m-onnx
|
| 37 |
+
results:
|
| 38 |
+
- task:
|
| 39 |
+
type: zero-shot-classification
|
| 40 |
+
name: Zero-shot intent routing
|
| 41 |
+
dataset:
|
| 42 |
+
type: custom
|
| 43 |
+
name: Held-out routing suite (1,636 messages, 10 workflows)
|
| 44 |
+
metrics:
|
| 45 |
+
- type: accuracy
|
| 46 |
+
value: 0.949
|
| 47 |
+
name: Routing accuracy at threshold 0.625
|
| 48 |
+
- type: recall
|
| 49 |
+
value: 0.986
|
| 50 |
+
name: Out-of-scope recall
|
| 51 |
+
---
|
| 52 |
+
|
| 53 |
+
# Laya Intent Router 150M (ONNX, int8)
|
| 54 |
+
|
| 55 |
+
**Zero-shot intent classification that runs on a CPU in about 60 ms.** You give it a message and a list of intents written in plain English. It tells you which intent the message belongs to, or that it belongs to none of them.
|
| 56 |
+
|
| 57 |
+
No training. No labelled data. You write the intents when you call it, and you can change them on every request.
|
| 58 |
+
|
| 59 |
+
```
|
| 60 |
+
"i want to cancel my order 88213" -> cancel_order 0.97
|
| 61 |
+
"hey where's my parcel, it's been 5 days" -> order_status 0.97
|
| 62 |
+
"ordered black shoes, got blue ones lol" -> wrong_item 0.97
|
| 63 |
+
"what's the weather in paris" -> none (0.98 none)
|
| 64 |
+
"qwewqeqw" -> none (0.98 none)
|
| 65 |
+
```
|
| 66 |
+
|
| 67 |
+
That last line is why I built this. The original Laya sent `qwewqeqw` to `order_not_received` with 0.73 confidence. This one puts 0.98 on none.
|
| 68 |
+
|
| 69 |
+
## Try it
|
| 70 |
+
|
| 71 |
+
```bash
|
| 72 |
+
pip install onnxruntime tokenizers numpy huggingface_hub
|
| 73 |
+
```
|
| 74 |
+
|
| 75 |
+
```python
|
| 76 |
+
from huggingface_hub import hf_hub_download
|
| 77 |
+
import importlib.util, sys
|
| 78 |
+
|
| 79 |
+
spec = importlib.util.spec_from_file_location(
|
| 80 |
+
"router", hf_hub_download("vrajnotviraj/laya-intent-router-150m-onnx", "router.py"))
|
| 81 |
+
router = importlib.util.module_from_spec(spec); spec.loader.exec_module(router)
|
| 82 |
+
|
| 83 |
+
r = router.Router.from_pretrained() # downloads ~440 MB once
|
| 84 |
+
|
| 85 |
+
intents = {
|
| 86 |
+
"check_balance": ["What's my account balance"],
|
| 87 |
+
"transfer_money": ["Send money to someone", "Transfer funds between accounts"],
|
| 88 |
+
"block_card": ["Block my card", "My card was lost or stolen"],
|
| 89 |
+
"loan_enquiry": ["Questions about personal or home loans"],
|
| 90 |
+
}
|
| 91 |
+
|
| 92 |
+
r.route("i think someone stole my card", intents)
|
| 93 |
+
# {'match': 'block_card', 'score': 0.98, 'probabilities': {..., '__none__': 0.01}}
|
| 94 |
+
|
| 95 |
+
r.route("ok", intents)
|
| 96 |
+
# {'match': None, 'score': 0.0, 'probabilities': {..., '__none__': 0.98}}
|
| 97 |
+
```
|
| 98 |
+
|
| 99 |
+
Each intent is a key plus 1 or more example phrasings or a short description. `match` is `None` when the message fits nothing, or when the best score is under the threshold (0.625 by default; pass `threshold=` to change it).
|
| 100 |
+
|
| 101 |
+
Or just clone the repo and run `python router.py "your message here"`.
|
| 102 |
+
|
| 103 |
+
## More examples
|
| 104 |
+
|
| 105 |
+
A SaaS support bot, 4 intents, nothing fine-tuned:
|
| 106 |
+
|
| 107 |
+
| message | match | score |
|
| 108 |
+
|---|---|---|
|
| 109 |
+
| cant get into my account | reset_password | 0.95 |
|
| 110 |
+
| why was i charged twice this month | billing | 0.98 |
|
| 111 |
+
| the export button does nothing | bug_report | 0.73 |
|
| 112 |
+
| can i talk to an actual human please | talk_to_human | 0.99 |
|
| 113 |
+
|
| 114 |
+
The intents were `reset_password: "User can't log in or forgot their password"`, `billing: "Questions about invoices, charges or refunds"`, `bug_report: "Something in the app is broken or showing an error"` and `talk_to_human: "User wants to speak to a real person"`. That's the whole setup.
|
| 115 |
+
|
| 116 |
+
## How good is it
|
| 117 |
+
|
| 118 |
+
I tested it on 1,636 hand-written messages across 10 routing setups: e-commerce, retail banking, insurance and telecom, plus an adversarial set of near-duplicates, typos, slang and "don't cancel, just tell me where it is" style traps. **Insurance and telecom were never seen in training.**
|
| 119 |
+
|
| 120 |
+
| | Laya (original, zero-shot) | Laya-large fine-tuned (teacher) | **This model** |
|
| 121 |
+
|---|---|---|---|
|
| 122 |
+
| Routing accuracy | 0.728 | 0.950 | **0.949** |
|
| 123 |
+
| Catches out-of-scope messages | 0.682 | 0.959 | **0.986** |
|
| 124 |
+
| Wrongly accepts out-of-scope | 0.318 | 0.041 | **0.014** |
|
| 125 |
+
| Big menus (20 to 45 intents) | 0.704 | 0.938 | **0.943** |
|
| 126 |
+
| 6 brand new domains | 0.777 | | **0.920** |
|
| 127 |
+
| 166 extra test messages, written last | 0.566 | 0.934 | **0.952** |
|
| 128 |
+
| p95 latency, 4 CPU threads | 252 ms | 496 ms | **160 ms** |
|
| 129 |
+
| Size on disk | 636 MB | 636 MB | **304 MB** (+134 MB shortlist) |
|
| 130 |
+
|
| 131 |
+
So you get the big fine-tuned model's accuracy at a third of its latency and half its size. Latency was measured on an Apple M2 Pro. Your server's vCPUs are probably slower, so measure there.
|
| 132 |
+
|
| 133 |
+
It also holds up when I poke at it. A perturbation test (shuffled intent order, removed correct intent, distractor intents, rewritten messages, gibberish) scores 0.92 averaged over 3 seeds, against 0.73 for the original Laya.
|
| 134 |
+
|
| 135 |
+
## Big intent lists
|
| 136 |
+
|
| 137 |
+
The model reads everything in one 512-token window, so it can only see so many intents at once. For lists longer than 4, a small embedder (`bge-small-en-v1.5`, bundled in `shortlist/`) picks the 4 closest intents first and the router decides between those and "none".
|
| 138 |
+
|
| 139 |
+
That's on by default and it's why accuracy stays flat as the list grows:
|
| 140 |
+
|
| 141 |
+
| intents | full list | with shortlist |
|
| 142 |
+
|---|---|---|
|
| 143 |
+
| 7 | 0.915 | 0.930 |
|
| 144 |
+
| 12 | 0.969 | 0.979 |
|
| 145 |
+
| 45 | 0.781 | 0.938 |
|
| 146 |
+
| 148 | too long to run | 0.889 |
|
| 147 |
+
|
| 148 |
+
The first call on a new intent list embeds all its phrasings (about 0.9 s for 50 intents), then it's cached. Later calls stay around 180 ms p95 whatever the list size. Pass `shortlist_k=0` to turn it off.
|
| 149 |
+
|
| 150 |
+
## How it was made
|
| 151 |
+
|
| 152 |
+
The base is [Laya](https://huggingface.co/convaiinnovations/laya), a ModernBERT-large decision model that scores a list of options in one forward pass. Out of the box it was too eager to match, so:
|
| 153 |
+
|
| 154 |
+
1. I fine-tuned Laya-large as a teacher on about 122k routing episodes built from 7 public intent datasets (CLINC150, BANKING77, HWU64, SNIPS and 3 Bitext customer-support sets) plus 76 synthetic workflows across 16 domains. About 40% of episodes had the right intent removed, so the model learns to say "none". Loss was soft cross-entropy plus Laya's RL objective.
|
| 155 |
+
2. I rewrote the prompt format. Each intent went from a wordy wrapper to `key: "phrasing 1" | "phrasing 2"`, with leftover token budget handed to intents that need it. That alone was worth about 3 points.
|
| 156 |
+
3. I distilled the teacher into [Ettin-150M](https://huggingface.co/jhu-clsp/ettin-encoder-150m) with a fresh Laya head, mixing the teacher's probabilities 50/50 with the gold label. Checkpoints were picked by a perturbation-based intent score, since plain accuracy picked worse routers.
|
| 157 |
+
4. I calibrated temperatures per menu size, then exported to ONNX with 8-bit weight quantization.
|
| 158 |
+
|
| 159 |
+
Everything trained locally on a 32 GB M2 Pro. Training code: [github.com/vrajnotviraj/laya-intent-router](https://github.com/vrajnotviraj/laya-intent-router).
|
| 160 |
+
|
| 161 |
+
## Where it slips
|
| 162 |
+
|
| 163 |
+
- **English only.** It hasn't seen other languages.
|
| 164 |
+
- **Filler on tiny menus.** With only 2 intents like "Hi" and "Bye", words like "ok", "well" and "nice" can land on "Bye" (0.67 to 0.81). Give short closing intents a clear description, or raise the threshold for small menus.
|
| 165 |
+
- **It only knows what your phrasings say.** "My card was stolen" won't hit a `block_card` intent whose only phrasing is "Block my card". Add 2 or 3 phrasings that cover how people actually talk.
|
| 166 |
+
- **Indirect requests and heavy typos** are the weakest slices (about 0.83 to 0.86).
|
| 167 |
+
- My test messages were written by the same process as the synthetic training data. Real user logs might be harder. Treat the numbers as a strong hint and run your own messages through it.
|
| 168 |
+
|
| 169 |
+
## Files
|
| 170 |
+
|
| 171 |
+
| file | what |
|
| 172 |
+
|---|---|
|
| 173 |
+
| `laya.onnx`, `laya.onnx.data` | the router (int8 weights) |
|
| 174 |
+
| `tokenizer.json`, `rl_agent_config.json` | tokenizer, prompt format, calibrated temperatures, threshold |
|
| 175 |
+
| `shortlist/` | bge-small-en-v1.5 ONNX embedder for long intent lists |
|
| 176 |
+
| `router.py` | the whole inference code, one file, no torch |
|
| 177 |
+
|
| 178 |
+
## License and credits
|
| 179 |
+
|
| 180 |
+
Apache-2.0. Built on [convaiinnovations/laya](https://huggingface.co/convaiinnovations/laya) (Apache-2.0, ModernBERT-large), [jhu-clsp/ettin-encoder-150m](https://huggingface.co/jhu-clsp/ettin-encoder-150m) (MIT) and [BAAI/bge-small-en-v1.5](https://huggingface.co/BAAI/bge-small-en-v1.5) (MIT).
|
| 181 |
+
|
| 182 |
+
Training data: CLINC150 (CC BY 3.0, Larson et al. 2019), BANKING77 (CC BY 4.0, Casanueva et al. 2020), HWU64 (CC BY 4.0, Liu et al. 2019), SNIPS (CC0) and the Bitext customer-support, retail e-commerce and retail banking datasets (CDLA-Sharing-1.0). No training data is included here.
|
laya.onnx
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:53e526532f19fff929506f3bffbd4faa0f6fdb684fa5842a2376e973e4091808
|
| 3 |
+
size 2467583
|
laya.onnx.data
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:17cfafb27b977cd0537a61da24edecad8d8f50eac2d0eb9f3e14e5bc3116cf9c
|
| 3 |
+
size 299825152
|
rl_agent_config.json
ADDED
|
@@ -0,0 +1,27 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"head_layers": 2,
|
| 3 |
+
"max_len": 512,
|
| 4 |
+
"head_max_len": 384,
|
| 5 |
+
"max_prefixes": 6,
|
| 6 |
+
"act_costs": {
|
| 7 |
+
"escalate": 0.5
|
| 8 |
+
},
|
| 9 |
+
"cost_wrong_act": 3.0,
|
| 10 |
+
"encoder": "jhu-clsp/ettin-encoder-150m",
|
| 11 |
+
"model_name": "laya-student",
|
| 12 |
+
"amp_dtype": "bf16",
|
| 13 |
+
"temperature": [
|
| 14 |
+
1.1014399528503418,
|
| 15 |
+
1.0,
|
| 16 |
+
1.0
|
| 17 |
+
],
|
| 18 |
+
"option_format": "compact_fill",
|
| 19 |
+
"temperature_by_options": {
|
| 20 |
+
"choice:2": 1.1014399528503418,
|
| 21 |
+
"choice:3-5": 1.1109503507614136,
|
| 22 |
+
"choice:6-10": 1.1113102436065674,
|
| 23 |
+
"choice:11+": 1.0
|
| 24 |
+
},
|
| 25 |
+
"match_threshold": 0.625,
|
| 26 |
+
"model_id": "vrajnotviraj/laya-intent-router-150m-onnx"
|
| 27 |
+
}
|
router.py
ADDED
|
@@ -0,0 +1,203 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Zero-shot intent router: pick which of your paths a message belongs to, or none of them.
|
| 2 |
+
|
| 3 |
+
from router import Router
|
| 4 |
+
r = Router.from_pretrained() # or Router("path/to/local/dir")
|
| 5 |
+
r.route("i want to cancel my order 88213", {
|
| 6 |
+
"order_status": ["Customer wants to know where their order is"],
|
| 7 |
+
"cancel_order": ["Customer wants to cancel an existing order"],
|
| 8 |
+
})
|
| 9 |
+
# {'match': 'cancel_order', 'score': 0.97, 'probabilities': {...}}
|
| 10 |
+
|
| 11 |
+
Needs: pip install onnxruntime tokenizers numpy huggingface_hub
|
| 12 |
+
"""
|
| 13 |
+
|
| 14 |
+
import hashlib
|
| 15 |
+
import json
|
| 16 |
+
import os
|
| 17 |
+
import re
|
| 18 |
+
from collections import OrderedDict
|
| 19 |
+
|
| 20 |
+
import numpy as np
|
| 21 |
+
|
| 22 |
+
REPO_ID = "vrajnotviraj/laya-intent-router-150m-onnx"
|
| 23 |
+
INSTRUCTIONS = (
|
| 24 |
+
"A user sent this message to a conversational workflow that branches into the paths below. "
|
| 25 |
+
"Which path is the message asking for?"
|
| 26 |
+
)
|
| 27 |
+
HISTORY_HINT = (
|
| 28 |
+
" `history` holds the earlier turns of this conversation, oldest first; a short or elliptical "
|
| 29 |
+
"message usually continues the intent of the most recent turns."
|
| 30 |
+
)
|
| 31 |
+
NO_MATCH_DESCRIPTION = "Gibberish, filler words, or a message unrelated to every other path."
|
| 32 |
+
NONE = "__none__"
|
| 33 |
+
OPTION_CAP, MIN_HEAD_TOKENS, PHRASING_SEP = 96, 16, '" | "'
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
def _water_fill(lengths, budget, floor=4):
|
| 37 |
+
open_ = list(range(len(lengths)))
|
| 38 |
+
while open_:
|
| 39 |
+
fair = budget // len(open_)
|
| 40 |
+
short = [i for i in open_ if lengths[i] <= fair]
|
| 41 |
+
if not short:
|
| 42 |
+
break
|
| 43 |
+
budget -= sum(lengths[i] for i in short)
|
| 44 |
+
open_ = [i for i in open_ if lengths[i] > fair]
|
| 45 |
+
alloc = list(lengths)
|
| 46 |
+
for j, i in enumerate(open_):
|
| 47 |
+
alloc[i] = max(floor, budget // len(open_) + (j < budget % len(open_)))
|
| 48 |
+
return alloc
|
| 49 |
+
|
| 50 |
+
|
| 51 |
+
def _humanise(path_id):
|
| 52 |
+
s = re.sub(r"([a-z0-9])([A-Z])", r"\1 \2", path_id)
|
| 53 |
+
return re.sub(r"[_\-.:/]+", " ", s).strip().lower()
|
| 54 |
+
|
| 55 |
+
|
| 56 |
+
def _session(path, threads):
|
| 57 |
+
import onnxruntime as ort
|
| 58 |
+
|
| 59 |
+
o = ort.SessionOptions()
|
| 60 |
+
o.intra_op_num_threads, o.inter_op_num_threads = threads, 1
|
| 61 |
+
return ort.InferenceSession(path, o, providers=["CPUExecutionProvider"])
|
| 62 |
+
|
| 63 |
+
|
| 64 |
+
class Shortlist:
|
| 65 |
+
"""bge-small embedder: keeps the k paths closest to the message (best phrasing wins). Cached per path set."""
|
| 66 |
+
|
| 67 |
+
def __init__(self, model_dir, threads=4, cache_size=512):
|
| 68 |
+
from tokenizers import Tokenizer
|
| 69 |
+
|
| 70 |
+
self.tok = Tokenizer.from_file(os.path.join(model_dir, "tokenizer.json"))
|
| 71 |
+
self.tok.enable_truncation(128)
|
| 72 |
+
self.tok.enable_padding(pad_id=self.tok.token_to_id("[PAD]") or 0)
|
| 73 |
+
self.sess = _session(os.path.join(model_dir, "model.onnx"), threads)
|
| 74 |
+
self.cache, self.cache_size = OrderedDict(), cache_size
|
| 75 |
+
|
| 76 |
+
def _embed(self, texts):
|
| 77 |
+
encs = self.tok.encode_batch(texts)
|
| 78 |
+
ids = np.array([e.ids for e in encs], dtype=np.int64)
|
| 79 |
+
mask = np.array([e.attention_mask for e in encs], dtype=np.int64)
|
| 80 |
+
v = self.sess.run(None, {"input_ids": ids, "attention_mask": mask, "token_type_ids": np.zeros_like(ids)})[0][:, 0]
|
| 81 |
+
return v / np.maximum(np.linalg.norm(v, axis=1, keepdims=True), 1e-8)
|
| 82 |
+
|
| 83 |
+
def top(self, message, paths, k):
|
| 84 |
+
key = hashlib.blake2b(json.dumps(paths).encode(), digest_size=16).hexdigest()
|
| 85 |
+
if key not in self.cache:
|
| 86 |
+
texts, owner = [], []
|
| 87 |
+
for j, texts_j in enumerate(paths.values()):
|
| 88 |
+
for t in list(texts_j) + [_humanise(list(paths)[j])]:
|
| 89 |
+
texts.append(t)
|
| 90 |
+
owner.append(j)
|
| 91 |
+
self.cache[key] = (self._embed(texts), np.array(owner))
|
| 92 |
+
if len(self.cache) > self.cache_size:
|
| 93 |
+
self.cache.popitem(last=False)
|
| 94 |
+
self.cache.move_to_end(key)
|
| 95 |
+
vecs, owner = self.cache[key]
|
| 96 |
+
sims = vecs @ self._embed([message])[0]
|
| 97 |
+
best = np.array([sims[owner == j].max() for j in range(len(paths))])
|
| 98 |
+
keep = set(np.argsort(-best, kind="stable")[:k].tolist())
|
| 99 |
+
return {p: t for j, (p, t) in enumerate(paths.items()) if j in keep}
|
| 100 |
+
|
| 101 |
+
|
| 102 |
+
class Router:
|
| 103 |
+
def __init__(self, model_dir, threads=4, shortlist_k=4):
|
| 104 |
+
from tokenizers import Tokenizer
|
| 105 |
+
|
| 106 |
+
self.cfg = json.load(open(os.path.join(model_dir, "rl_agent_config.json")))
|
| 107 |
+
self.tok = Tokenizer.from_file(os.path.join(model_dir, "tokenizer.json"))
|
| 108 |
+
self.cls, self.sep, self.mask = (self.tok.token_to_id(t) for t in ("[CLS]", "[SEP]", "[MASK]"))
|
| 109 |
+
self.sess = _session(os.path.join(model_dir, "laya.onnx"), threads)
|
| 110 |
+
self.threshold = self.cfg.get("match_threshold", 0.625)
|
| 111 |
+
self.temps = self.cfg["temperature_by_options"]
|
| 112 |
+
sl_dir = os.path.join(model_dir, "shortlist")
|
| 113 |
+
self.k = shortlist_k if os.path.isdir(sl_dir) else 0
|
| 114 |
+
self.shortlist = Shortlist(sl_dir, threads) if self.k else None
|
| 115 |
+
|
| 116 |
+
@classmethod
|
| 117 |
+
def from_pretrained(cls, repo_id=REPO_ID, **kw):
|
| 118 |
+
from huggingface_hub import snapshot_download
|
| 119 |
+
|
| 120 |
+
return cls(snapshot_download(repo_id), **kw)
|
| 121 |
+
|
| 122 |
+
def _tokens(self, text):
|
| 123 |
+
return self.tok.encode(text.replace("[MASK]", " "), add_special_tokens=False)
|
| 124 |
+
|
| 125 |
+
def _options(self, criteria, head_max_len):
|
| 126 |
+
texts = [f" {k}: {d}".replace("[MASK]", " ") for k, d in criteria.items()]
|
| 127 |
+
encs = [self.tok.encode(t, add_special_tokens=False) for t in texts]
|
| 128 |
+
options = [[self.mask] + e.ids[:OPTION_CAP] for e in encs]
|
| 129 |
+
if head_max_len - sum(map(len, options)) < MIN_HEAD_TOKENS: # too many paths: share the budget fairly
|
| 130 |
+
alloc = _water_fill([len(o) for o in options], head_max_len - MIN_HEAD_TOKENS)
|
| 131 |
+
for i, (t, e, a) in enumerate(zip(texts, encs, alloc)):
|
| 132 |
+
n = a - 1
|
| 133 |
+
if n >= len(options[i]) - 1:
|
| 134 |
+
continue
|
| 135 |
+
ends, keep = [end for _, end in e.offsets], n
|
| 136 |
+
p = t.find(PHRASING_SEP)
|
| 137 |
+
while p >= 0: # prefer cutting between phrasings
|
| 138 |
+
b = sum(x <= p + 1 for x in ends)
|
| 139 |
+
if b > n:
|
| 140 |
+
break
|
| 141 |
+
if 4 * (n - b) <= a:
|
| 142 |
+
keep = b
|
| 143 |
+
p = t.find(PHRASING_SEP, p + 1)
|
| 144 |
+
options[i] = [self.mask] + e.ids[:keep]
|
| 145 |
+
return options
|
| 146 |
+
|
| 147 |
+
def _encode(self, message, criteria, history):
|
| 148 |
+
max_len, head_max_len = self.cfg["max_len"], self.cfg["head_max_len"]
|
| 149 |
+
options = self._options(criteria, head_max_len)
|
| 150 |
+
instructions = INSTRUCTIONS + (HISTORY_HINT if history else "")
|
| 151 |
+
head = self._tokens(f"choice question: {instructions}").ids[: max(8, head_max_len - sum(map(len, options)))]
|
| 152 |
+
ids, markers = [self.cls, *head, self.sep], []
|
| 153 |
+
for o in options:
|
| 154 |
+
markers.append(len(ids))
|
| 155 |
+
ids += o
|
| 156 |
+
ids.append(self.sep)
|
| 157 |
+
state = {"message": message, **({"history": list(history)} if history else {})}
|
| 158 |
+
room = max(0, max_len - len(ids) - 1)
|
| 159 |
+
ids = (ids + self._tokens(json.dumps(state, ensure_ascii=False)).ids[:room] + [self.sep])[:max_len]
|
| 160 |
+
if any(m >= max_len for m in markers):
|
| 161 |
+
raise ValueError("too many paths for 512 tokens; keep the shortlist on")
|
| 162 |
+
return ids, markers
|
| 163 |
+
|
| 164 |
+
def _temperature(self, n):
|
| 165 |
+
b = "2" if n <= 2 else "3-5" if n <= 5 else "6-10" if n <= 10 else "11+"
|
| 166 |
+
return min(5.0, max(0.5, float(self.temps.get(f"choice:{b}", self.cfg["temperature"][0]))))
|
| 167 |
+
|
| 168 |
+
def route(self, message, paths, history=None, threshold=None):
|
| 169 |
+
"""paths: {path_id: [example phrasings or a description, ...]}. Returns the match (or None) and every probability."""
|
| 170 |
+
none_key = NONE
|
| 171 |
+
while none_key in paths:
|
| 172 |
+
none_key += "_"
|
| 173 |
+
kept = self.shortlist.top(message, paths, self.k) if self.k and len(paths) > self.k else paths
|
| 174 |
+
criteria = {p: " | ".join(f'"{t}"' for t in texts) for p, texts in kept.items()}
|
| 175 |
+
criteria[none_key] = NO_MATCH_DESCRIPTION
|
| 176 |
+
ids, markers = self._encode(message, criteria, history)
|
| 177 |
+
logits = self.sess.run(["logits"], {
|
| 178 |
+
"input_ids": np.array([ids], dtype=np.int64), "attention_mask": np.ones((1, len(ids)), dtype=np.int64),
|
| 179 |
+
"marker_pos": np.array([markers], dtype=np.int64), "marker_mask": np.ones((1, len(markers)), dtype=bool),
|
| 180 |
+
"qtype": np.zeros(1, dtype=np.int64)})[0][0]
|
| 181 |
+
z = logits / self._temperature(len(criteria))
|
| 182 |
+
p = np.exp(z - z.max())
|
| 183 |
+
p /= p.sum()
|
| 184 |
+
got = dict(zip(criteria, p.tolist()))
|
| 185 |
+
probs = {k: got.get(k, 0.0) for k in paths} | {NONE: got[none_key]}
|
| 186 |
+
best = max(paths, key=probs.get)
|
| 187 |
+
ok = probs[best] > probs[NONE] and probs[best] >= (self.threshold if threshold is None else threshold)
|
| 188 |
+
return {"match": best if ok else None, "score": probs[best], "probabilities": probs}
|
| 189 |
+
|
| 190 |
+
|
| 191 |
+
if __name__ == "__main__":
|
| 192 |
+
import sys
|
| 193 |
+
|
| 194 |
+
r = Router(os.path.dirname(os.path.abspath(__file__)))
|
| 195 |
+
paths = {
|
| 196 |
+
"order_status": ["Customer wants to know the status or location of their order"],
|
| 197 |
+
"cancel_order": ["Customer wants to cancel an existing order"],
|
| 198 |
+
"return_order": ["Customer wants to return an order"],
|
| 199 |
+
"wrong_item": ["Customer received an item different from what they ordered"],
|
| 200 |
+
}
|
| 201 |
+
for msg in sys.argv[1:] or ["where is my parcel #A-7721", "qwewqeqw", "i got blue shoes but ordered black"]:
|
| 202 |
+
out = r.route(msg, paths)
|
| 203 |
+
print(f"{msg!r:45} -> {out['match']} ({out['score']:.2f})")
|
shortlist/model.onnx
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:828e1496d7fabb79cfa4dcd84fa38625c0d3d21da474a00f08db0f559940cf35
|
| 3 |
+
size 133093490
|
shortlist/shortlist_config.json
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
{"pooling": "cls", "max_len": 128}
|
shortlist/tokenizer.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
tokenizer.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|