lokinfey commited on
Commit
18a6fb5
·
verified ·
1 Parent(s): f987391

Add typed-decision ONNX runner

Browse files
Files changed (1) hide show
  1. run_hmm_onnx.py +364 -0
run_hmm_onnx.py ADDED
@@ -0,0 +1,364 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import argparse
4
+ import json
5
+ import math
6
+ import os
7
+ import time
8
+ from pathlib import Path
9
+ from typing import Any
10
+
11
+ import numpy as np
12
+ import onnxruntime as ort
13
+ from transformers import AutoTokenizer
14
+
15
+
16
+ LETTERS = "ABCDEFGHIJKLMNOPQRSTUVWXYZ"
17
+ MAX_OPTIONS = 255
18
+ TOKENIZER_REPO = "Qwen/Qwen3.5-4B"
19
+ TOKENIZER_REVISION = "851bf6e806efd8d0a36b00ddf55e13ccb7b8cd0a"
20
+
21
+ ORT_TO_NUMPY = {
22
+ "tensor(float)": np.float32,
23
+ "tensor(float16)": np.float16,
24
+ "tensor(double)": np.float64,
25
+ "tensor(int64)": np.int64,
26
+ "tensor(int32)": np.int32,
27
+ "tensor(bool)": np.bool_,
28
+ }
29
+
30
+
31
+ def text(value: Any) -> str:
32
+ if isinstance(value, str):
33
+ return value
34
+ return json.dumps(value, ensure_ascii=False, separators=(",", ":"))
35
+
36
+
37
+ def options_for(name: str, question: dict[str, Any]) -> list[tuple[str, str]]:
38
+ if not isinstance(question, dict) or question.get("instructions") is None:
39
+ raise ValueError(f'question "{name}" needs "type" and "instructions"')
40
+
41
+ question_type = question.get("type")
42
+ criteria = question.get("criteria")
43
+
44
+ if question_type == "choice":
45
+ if not isinstance(criteria, dict) or not 2 <= len(criteria) <= MAX_OPTIONS:
46
+ raise ValueError(f'choice "{name}" needs 2-{MAX_OPTIONS} options')
47
+ return [(str(key), text(value)) for key, value in criteria.items()]
48
+
49
+ if question_type == "score":
50
+ if not isinstance(criteria, list) or not 2 <= len(criteria) <= 10:
51
+ raise ValueError(f'score "{name}" needs 2-10 ordered levels')
52
+ return [(str(index), text(value)) for index, value in enumerate(criteria)]
53
+
54
+ if question_type == "noul":
55
+ criteria = criteria if isinstance(criteria, dict) else {}
56
+ return [
57
+ ("false", text(criteria.get("false", "No"))),
58
+ ("true", text(criteria.get("true", "Yes"))),
59
+ ]
60
+
61
+ raise ValueError(f'question "{name}" has unknown type {question_type!r}')
62
+
63
+
64
+ def build_prompt(
65
+ state: Any,
66
+ question: dict[str, Any],
67
+ options: list[tuple[str, str]],
68
+ ) -> str:
69
+ lines = [
70
+ f"{LETTERS[index]}: {key} — {description}"
71
+ for index, (key, description) in enumerate(options)
72
+ ]
73
+ user = (
74
+ "State (data to evaluate):\n"
75
+ + text(state)
76
+ + "\n\nQuestion:\n"
77
+ + text(question["instructions"])
78
+ + "\n\nOptions:\n"
79
+ + "\n".join(lines)
80
+ + "\nReturn only the option letter."
81
+ )
82
+ return (
83
+ f"<|im_start|>user\n{user}<|im_end|>\n"
84
+ "<|im_start|>assistant\n<think>\n\n</think>\n\n"
85
+ )
86
+
87
+
88
+ class HmmOnnx:
89
+ def __init__(self, model_dir: Path) -> None:
90
+ model_path = model_dir / "model.onnx"
91
+ if not model_path.is_file():
92
+ raise FileNotFoundError(model_path)
93
+
94
+ options = ort.SessionOptions()
95
+ options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
96
+ options.enable_mem_pattern = False
97
+ options.intra_op_num_threads = max(1, (os.cpu_count() or 2) // 2)
98
+
99
+ self.session = ort.InferenceSession(
100
+ str(model_path),
101
+ sess_options=options,
102
+ providers=["CPUExecutionProvider"],
103
+ )
104
+ self.input_names = {item.name for item in self.session.get_inputs()}
105
+ self.output_names = [item.name for item in self.session.get_outputs()]
106
+ self.tokenizer = AutoTokenizer.from_pretrained(
107
+ model_dir,
108
+ local_files_only=True,
109
+ trust_remote_code=False,
110
+ )
111
+ self.letter_token_ids = {}
112
+ for letter in LETTERS:
113
+ token_ids = self.tokenizer.encode(letter, add_special_tokens=False)
114
+ if len(token_ids) != 1:
115
+ raise ValueError(f"Option letter {letter!r} is not one token: {token_ids}")
116
+ self.letter_token_ids[letter] = token_ids[0]
117
+
118
+ @staticmethod
119
+ def _state_shape(model_input: Any) -> tuple[int, ...]:
120
+ shape = []
121
+ for axis, dimension in enumerate(model_input.shape):
122
+ if isinstance(dimension, int):
123
+ shape.append(dimension)
124
+ elif (
125
+ model_input.name.endswith((".key", ".value"))
126
+ and axis == 2
127
+ ):
128
+ shape.append(0)
129
+ else:
130
+ shape.append(1)
131
+ return tuple(shape)
132
+
133
+ def _initial_states(self) -> dict[str, np.ndarray]:
134
+ states = {}
135
+ for model_input in self.session.get_inputs():
136
+ if not model_input.name.startswith("past_key_values."):
137
+ continue
138
+ dtype = ORT_TO_NUMPY.get(model_input.type)
139
+ if dtype is None:
140
+ raise TypeError(f"Unsupported state dtype: {model_input.type}")
141
+ states[model_input.name] = np.zeros(
142
+ self._state_shape(model_input),
143
+ dtype=dtype,
144
+ )
145
+ return states
146
+
147
+ def _forward_prompt(self, prompt: str) -> tuple[np.ndarray, int, float]:
148
+ input_ids = self.tokenizer(
149
+ prompt,
150
+ add_special_tokens=False,
151
+ return_tensors="np",
152
+ )["input_ids"].astype(np.int64)
153
+
154
+ states = self._initial_states()
155
+ outputs = None
156
+ started = time.perf_counter()
157
+
158
+ for position, token_id in enumerate(input_ids[0]):
159
+ feeds = {
160
+ "input_ids": np.asarray([[token_id]], dtype=np.int64),
161
+ "attention_mask": np.ones((1, position + 1), dtype=np.int64),
162
+ "position_ids": np.asarray([[position]], dtype=np.int64),
163
+ **states,
164
+ }
165
+ missing = self.input_names - feeds.keys()
166
+ if missing:
167
+ raise ValueError(f"Missing model inputs: {sorted(missing)}")
168
+
169
+ values = self.session.run(
170
+ self.output_names,
171
+ {name: feeds[name] for name in self.input_names},
172
+ )
173
+ outputs = dict(zip(self.output_names, values, strict=True))
174
+ states = {
175
+ name: outputs[
176
+ f"present.{name.removeprefix('past_key_values.')}"
177
+ ]
178
+ for name in states
179
+ }
180
+
181
+ if outputs is None:
182
+ raise ValueError("Prompt produced no input tokens")
183
+
184
+ logits = outputs["logits"][0, -1].astype(np.float64)
185
+ if not np.isfinite(logits).all():
186
+ raise FloatingPointError("Non-finite logits detected")
187
+
188
+ return logits, int(input_ids.shape[1]), time.perf_counter() - started
189
+
190
+ def _letter_probabilities(
191
+ self,
192
+ prompt: str,
193
+ letters: str,
194
+ ) -> tuple[dict[str, float], int, float]:
195
+ logits, input_tokens, elapsed = self._forward_prompt(prompt)
196
+ shifted = logits - np.max(logits)
197
+ denominator = float(np.exp(shifted).sum())
198
+ probabilities = {
199
+ letter: float(
200
+ np.exp(shifted[self.letter_token_ids[letter]]) / denominator
201
+ )
202
+ for letter in letters
203
+ }
204
+ return probabilities, input_tokens, elapsed
205
+
206
+ @staticmethod
207
+ def _round4(value: float) -> float:
208
+ return math.floor(value * 10000 + 0.5) / 10000
209
+
210
+ def answer(
211
+ self,
212
+ state: Any,
213
+ name: str,
214
+ question: dict[str, Any],
215
+ ) -> dict[str, Any]:
216
+ options = options_for(name, question)
217
+ raw_probabilities = []
218
+ input_tokens = 0
219
+ elapsed = 0.0
220
+
221
+ if len(options) <= len(LETTERS):
222
+ letters = LETTERS[: len(options)]
223
+ raw, count, seconds = self._letter_probabilities(
224
+ build_prompt(state, question, options),
225
+ letters,
226
+ )
227
+ raw_probabilities = [raw[letter] for letter in letters]
228
+ input_tokens += count
229
+ elapsed += seconds
230
+ else:
231
+ chunk_size = len(LETTERS) - 1
232
+ none_option = ("none_of_these", "None of the other options fits")
233
+ for start in range(0, len(options), chunk_size):
234
+ chunk = options[start : start + chunk_size]
235
+ chunk_options = [*chunk, none_option]
236
+ letters = LETTERS[: len(chunk_options)]
237
+ raw, count, seconds = self._letter_probabilities(
238
+ build_prompt(state, question, chunk_options),
239
+ letters,
240
+ )
241
+ raw_probabilities.extend(
242
+ raw[LETTERS[index]] for index in range(len(chunk))
243
+ )
244
+ input_tokens += count
245
+ elapsed += seconds
246
+
247
+ total = sum(raw_probabilities)
248
+ probabilities = (
249
+ [value / total for value in raw_probabilities]
250
+ if total > 0
251
+ else [1.0 / len(options)] * len(options)
252
+ )
253
+ keys = [key for key, _ in options]
254
+ best = int(np.argmax(probabilities))
255
+ probability_map = {
256
+ key: self._round4(probabilities[index])
257
+ for index, key in enumerate(keys)
258
+ }
259
+ confidence = self._round4(
260
+ (len(options) * probabilities[best] - 1) / (len(options) - 1)
261
+ )
262
+
263
+ question_type = question["type"]
264
+ if question_type == "noul":
265
+ result = {
266
+ "type": "noul",
267
+ "noul": self._round4(probabilities[1]),
268
+ }
269
+ elif question_type == "choice":
270
+ result = {
271
+ "type": "choice",
272
+ "choice": keys[best],
273
+ "probabilities": probability_map,
274
+ "confidence": confidence,
275
+ }
276
+ else:
277
+ result = {
278
+ "type": "score",
279
+ "score": self._round4(
280
+ sum(index * probability for index, probability in enumerate(probabilities))
281
+ ),
282
+ "legend": dict(options),
283
+ "probabilities": probability_map,
284
+ "confidence": confidence,
285
+ }
286
+
287
+ return {
288
+ "input_tokens": input_tokens,
289
+ "seconds": elapsed,
290
+ "result": result,
291
+ }
292
+
293
+ def decide(self, body: dict[str, Any]) -> dict[str, Any]:
294
+ if body.get("state") in (None, ""):
295
+ raise ValueError('"state" is required')
296
+ questions = body.get("questions")
297
+ if not isinstance(questions, dict) or not questions:
298
+ raise ValueError('"questions" must be a non-empty object')
299
+
300
+ started = time.perf_counter()
301
+ answers = {}
302
+ input_tokens = 0
303
+ inference_seconds = 0.0
304
+
305
+ for name, question in questions.items():
306
+ answer = self.answer(body["state"], name, question)
307
+ answers[name] = answer["result"]
308
+ input_tokens += answer["input_tokens"]
309
+ inference_seconds += answer["seconds"]
310
+
311
+ return {
312
+ "model": "Qwen3.5-4B-Hmm-Q4_K_M.onnx",
313
+ "answers": answers,
314
+ "usage": {"input_tokens": input_tokens, "output_tokens": 0},
315
+ "latency_ms": round((time.perf_counter() - started) * 1000),
316
+ "inference_seconds": round(inference_seconds, 4),
317
+ }
318
+
319
+
320
+ DEFAULT_REQUEST = {
321
+ "state": "Help! My payouts have failed for 3 days. I need the money today.",
322
+ "questions": {
323
+ "is_urgent": {
324
+ "type": "noul",
325
+ "instructions": "Does this message convey urgency?",
326
+ },
327
+ "department": {
328
+ "type": "choice",
329
+ "instructions": "Which team should handle this?",
330
+ "criteria": {
331
+ "billing": "Payments, invoicing, refunds",
332
+ "technical": "Bugs, outages, integrations",
333
+ "sales": "Pricing, upgrades, new accounts",
334
+ },
335
+ },
336
+ },
337
+ }
338
+
339
+
340
+ def main() -> None:
341
+ parser = argparse.ArgumentParser()
342
+ parser.add_argument(
343
+ "--model-dir",
344
+ type=Path,
345
+ default=Path(".mobius_colab_run/onnx_outputs"),
346
+ )
347
+ parser.add_argument(
348
+ "--request",
349
+ type=Path,
350
+ help="Optional JSON request file. Uses the model-card example by default.",
351
+ )
352
+ args = parser.parse_args()
353
+
354
+ request = (
355
+ json.loads(args.request.read_text(encoding="utf-8"))
356
+ if args.request
357
+ else DEFAULT_REQUEST
358
+ )
359
+ model = HmmOnnx(args.model_dir)
360
+ print(json.dumps(model.decide(request), ensure_ascii=False, indent=2))
361
+
362
+
363
+ if __name__ == "__main__":
364
+ main()