vrajnotviraj commited on
Commit
d228ae1
·
verified ·
1 Parent(s): 79e219d

Laya Intent Router 150M: int8 ONNX, shortlist embedder, router.py, model card

Browse files
.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