Ruiruiz30 commited on
Commit
2ffc677
·
verified ·
1 Parent(s): a1f6a44

add JevBench recalibration and runtime temperature scaling

Browse files
README.md CHANGED
@@ -47,8 +47,10 @@ The published Jev-Omni H200 numbers are not transferable to this Mac mini. This
47
  - Upstream unified verification cases: 4/4 argmax decisions matched after 4-bit conversion.
48
  - Six simple red/blue/green circle and square image checks: 6/6 color decisions matched.
49
  - Maximum absolute probability difference on the four upstream text cases: 0.244 in this small check.
50
- - JevBench and DecisionBench were **not** re-run for this conversion.
51
- - Quantization changes the probability distribution; probabilities are not recalibrated here and must not be treated as calibrated confidence.
 
 
52
 
53
  ## Installation
54
 
@@ -69,6 +71,7 @@ pip install -r requirements.txt
69
  ```bash
70
  python -m omni_mlx.classifier \
71
  --model . \
 
72
  --image /path/to/frame.png \
73
  --state "A kart is approaching a right turn." \
74
  --question "Which steering action is best?" \
@@ -78,6 +81,14 @@ python -m omni_mlx.classifier \
78
 
79
  The classifier returns candidate probabilities and the selected option. It does not generate a free-form explanation. Use `--image-tokens 70` when the scene contains small or dense visual details.
80
 
 
 
 
 
 
 
 
 
81
  ## Local conversion code
82
 
83
  `omni_mlx/convert.py` contains the conversion path used for this release. The original unquantized checkpoint is not bundled here; it can be obtained from the upstream repository under its own license and terms.
 
47
  - Upstream unified verification cases: 4/4 argmax decisions matched after 4-bit conversion.
48
  - Six simple red/blue/green circle and square image checks: 6/6 color decisions matched.
49
  - Maximum absolute probability difference on the four upstream text cases: 0.244 in this small check.
50
+ - Public JevBench v1.2 public items (easy + original + hard, 231 decisions) were re-run locally. Raw accuracy was **87.88%**; state-macro and micro are identical because each public item is one state/question. Raw ECE-10 was **0.04497**.
51
+ - A single global temperature was fit on even source rows (116 items) and checked on odd rows (115 items): `T=1.11517`. On all 231 items, ECE-10 was **0.03069** after scaling; the held-out ECE was **0.06774** versus raw **0.06261**, so this is a published post-hoc calibration artifact, not a universal confidence guarantee.
52
+ - DecisionBench is being evaluated with the original full states; its long 16GB-Mac run is checkpointed under `artifacts/benchmarks/runs/` in the source project. It is not included in the model-card score until the complete run finishes.
53
+ - The runtime supports the published temperature file through `--calibration calibration.json`. Accuracy/argmax is unchanged by temperature scaling; only the returned probability distribution changes.
54
 
55
  ## Installation
56
 
 
71
  ```bash
72
  python -m omni_mlx.classifier \
73
  --model . \
74
+ --calibration calibration.json \
75
  --image /path/to/frame.png \
76
  --state "A kart is approaching a right turn." \
77
  --question "Which steering action is best?" \
 
81
 
82
  The classifier returns candidate probabilities and the selected option. It does not generate a free-form explanation. Use `--image-tokens 70` when the scene contains small or dense visual details.
83
 
84
+ Omit `--calibration` to inspect the raw quantized probabilities. The included calibration file was fit only on the public JevBench split described above; it is not trained on a user's game or on private benchmark items.
85
+
86
+ ## Benchmark artifacts
87
+
88
+ - [`benchmarks/jevbench-4bit-report.json`](benchmarks/jevbench-4bit-report.json) contains the raw and temperature-scaled aggregate metrics.
89
+ - [`calibration.json`](calibration.json) is the small runtime file consumed by `--calibration`.
90
+ - The benchmark runner and raw checkpoints remain in the source project so the long DecisionBench run can be resumed without putting the full benchmark text into this model repository.
91
+
92
  ## Local conversion code
93
 
94
  `omni_mlx/convert.py` contains the conversion path used for this release. The original unquantized checkpoint is not bundled here; it can be obtained from the upstream repository under its own license and terms.
benchmarks/jevbench-4bit-report.json ADDED
@@ -0,0 +1,433 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "method": "temperature_scaling",
3
+ "temperature": 1.11516790625,
4
+ "fit_questions": 116,
5
+ "held_out_questions": 115,
6
+ "split": "even-rows",
7
+ "split_description": "even source rows fit; odd source rows held out, separately within each supplied tier",
8
+ "source_results": "artifacts/benchmarks/runs/jevbench-4bit-raw.jsonl",
9
+ "raw": {
10
+ "fit": {
11
+ "questions": 116,
12
+ "errors": 0,
13
+ "scenarios": 116,
14
+ "micro_accuracy": 0.8706896551724138,
15
+ "state_macro_accuracy": 0.8706896551724138,
16
+ "ece_10": 0.049421241828079915,
17
+ "confidence_gap": 0.014750512766427026,
18
+ "mean_confidence": 0.8854401679388408,
19
+ "latency_ms": {
20
+ "median": 1326.6938750020927,
21
+ "p95": 31386.2294999999
22
+ },
23
+ "option_count": {
24
+ "min": 2,
25
+ "max": 6,
26
+ "mean": 3.6724137931034484
27
+ },
28
+ "calibration_bins": [
29
+ {
30
+ "lower": 0.3,
31
+ "upper": 0.4,
32
+ "count": 3,
33
+ "accuracy": 0.3333333333333333,
34
+ "confidence": 0.35904427369435626
35
+ },
36
+ {
37
+ "lower": 0.4,
38
+ "upper": 0.5,
39
+ "count": 2,
40
+ "accuracy": 0.0,
41
+ "confidence": 0.44150300323963165
42
+ },
43
+ {
44
+ "lower": 0.5,
45
+ "upper": 0.6,
46
+ "count": 6,
47
+ "accuracy": 0.6666666666666666,
48
+ "confidence": 0.5567008852958679
49
+ },
50
+ {
51
+ "lower": 0.6,
52
+ "upper": 0.7,
53
+ "count": 4,
54
+ "accuracy": 0.5,
55
+ "confidence": 0.6537451297044754
56
+ },
57
+ {
58
+ "lower": 0.7,
59
+ "upper": 0.8,
60
+ "count": 11,
61
+ "accuracy": 0.8181818181818182,
62
+ "confidence": 0.7394869869405573
63
+ },
64
+ {
65
+ "lower": 0.8,
66
+ "upper": 0.9,
67
+ "count": 12,
68
+ "accuracy": 0.6666666666666666,
69
+ "confidence": 0.8455702016750971
70
+ },
71
+ {
72
+ "lower": 0.9,
73
+ "upper": 1.0,
74
+ "count": 78,
75
+ "accuracy": 0.9871794871794872,
76
+ "confidence": 0.9809555839269589
77
+ }
78
+ ]
79
+ },
80
+ "held_out": {
81
+ "questions": 115,
82
+ "errors": 0,
83
+ "scenarios": 115,
84
+ "micro_accuracy": 0.8869565217391304,
85
+ "state_macro_accuracy": 0.8869565217391304,
86
+ "ece_10": 0.06260819901590761,
87
+ "confidence_gap": -0.0037480162537616435,
88
+ "mean_confidence": 0.8832085054853688,
89
+ "latency_ms": {
90
+ "median": 1331.1896250088466,
91
+ "p95": 31440.3575410106
92
+ },
93
+ "option_count": {
94
+ "min": 2,
95
+ "max": 6,
96
+ "mean": 3.6782608695652175
97
+ },
98
+ "calibration_bins": [
99
+ {
100
+ "lower": 0.3,
101
+ "upper": 0.4,
102
+ "count": 1,
103
+ "accuracy": 0.0,
104
+ "confidence": 0.3194894790649414
105
+ },
106
+ {
107
+ "lower": 0.4,
108
+ "upper": 0.5,
109
+ "count": 7,
110
+ "accuracy": 0.42857142857142855,
111
+ "confidence": 0.4288983941078186
112
+ },
113
+ {
114
+ "lower": 0.5,
115
+ "upper": 0.6,
116
+ "count": 5,
117
+ "accuracy": 1.0,
118
+ "confidence": 0.5324515223503112
119
+ },
120
+ {
121
+ "lower": 0.6,
122
+ "upper": 0.7,
123
+ "count": 4,
124
+ "accuracy": 1.0,
125
+ "confidence": 0.6305650025606155
126
+ },
127
+ {
128
+ "lower": 0.7,
129
+ "upper": 0.8,
130
+ "count": 7,
131
+ "accuracy": 0.7142857142857143,
132
+ "confidence": 0.7388283269745963
133
+ },
134
+ {
135
+ "lower": 0.8,
136
+ "upper": 0.9,
137
+ "count": 11,
138
+ "accuracy": 0.7272727272727273,
139
+ "confidence": 0.8498979048295454
140
+ },
141
+ {
142
+ "lower": 0.9,
143
+ "upper": 1.0,
144
+ "count": 80,
145
+ "accuracy": 0.9625,
146
+ "confidence": 0.9817750878632069
147
+ }
148
+ ]
149
+ },
150
+ "all": {
151
+ "questions": 231,
152
+ "errors": 0,
153
+ "scenarios": 231,
154
+ "micro_accuracy": 0.8787878787878788,
155
+ "state_macro_accuracy": 0.8787878787878788,
156
+ "ece_10": 0.04497108405286617,
157
+ "confidence_gap": 0.005541288362436947,
158
+ "mean_confidence": 0.8843291671503157,
159
+ "latency_ms": {
160
+ "median": 1331.1896250088466,
161
+ "p95": 31386.2294999999
162
+ },
163
+ "option_count": {
164
+ "min": 2,
165
+ "max": 6,
166
+ "mean": 3.675324675324675
167
+ },
168
+ "calibration_bins": [
169
+ {
170
+ "lower": 0.3,
171
+ "upper": 0.4,
172
+ "count": 4,
173
+ "accuracy": 0.25,
174
+ "confidence": 0.34915557503700256
175
+ },
176
+ {
177
+ "lower": 0.4,
178
+ "upper": 0.5,
179
+ "count": 9,
180
+ "accuracy": 0.3333333333333333,
181
+ "confidence": 0.4316994183593326
182
+ },
183
+ {
184
+ "lower": 0.5,
185
+ "upper": 0.6,
186
+ "count": 11,
187
+ "accuracy": 0.8181818181818182,
188
+ "confidence": 0.5456784475933422
189
+ },
190
+ {
191
+ "lower": 0.6,
192
+ "upper": 0.7,
193
+ "count": 8,
194
+ "accuracy": 0.75,
195
+ "confidence": 0.6421550661325455
196
+ },
197
+ {
198
+ "lower": 0.7,
199
+ "upper": 0.8,
200
+ "count": 18,
201
+ "accuracy": 0.7777777777777778,
202
+ "confidence": 0.7392308413982391
203
+ },
204
+ {
205
+ "lower": 0.8,
206
+ "upper": 0.9,
207
+ "count": 23,
208
+ "accuracy": 0.6956521739130435,
209
+ "confidence": 0.8476399727489637
210
+ },
211
+ {
212
+ "lower": 0.9,
213
+ "upper": 1.0,
214
+ "count": 158,
215
+ "accuracy": 0.9746835443037974,
216
+ "confidence": 0.9813705226288566
217
+ }
218
+ ]
219
+ }
220
+ },
221
+ "temperature_scaled": {
222
+ "fit": {
223
+ "questions": 116,
224
+ "errors": 0,
225
+ "scenarios": 116,
226
+ "micro_accuracy": 0.8706896551724138,
227
+ "state_macro_accuracy": 0.8706896551724138,
228
+ "ece_10": 0.017033647453133953,
229
+ "confidence_gap": 0.00023089893600058975,
230
+ "mean_confidence": 0.8709205541084144,
231
+ "latency_ms": {
232
+ "median": 1326.6938750020927,
233
+ "p95": 31386.2294999999
234
+ },
235
+ "option_count": {
236
+ "min": 2,
237
+ "max": 6,
238
+ "mean": 3.6724137931034484
239
+ },
240
+ "calibration_bins": [
241
+ {
242
+ "lower": 0.3,
243
+ "upper": 0.4,
244
+ "count": 3,
245
+ "accuracy": 0.3333333333333333,
246
+ "confidence": 0.34551506179739544
247
+ },
248
+ {
249
+ "lower": 0.4,
250
+ "upper": 0.5,
251
+ "count": 4,
252
+ "accuracy": 0.5,
253
+ "confidence": 0.4578202850496621
254
+ },
255
+ {
256
+ "lower": 0.5,
257
+ "upper": 0.6,
258
+ "count": 5,
259
+ "accuracy": 0.6,
260
+ "confidence": 0.5667678762837054
261
+ },
262
+ {
263
+ "lower": 0.6,
264
+ "upper": 0.7,
265
+ "count": 8,
266
+ "accuracy": 0.625,
267
+ "confidence": 0.665780870468869
268
+ },
269
+ {
270
+ "lower": 0.7,
271
+ "upper": 0.8,
272
+ "count": 10,
273
+ "accuracy": 0.7,
274
+ "confidence": 0.7471593237748635
275
+ },
276
+ {
277
+ "lower": 0.8,
278
+ "upper": 0.9,
279
+ "count": 12,
280
+ "accuracy": 0.8333333333333334,
281
+ "confidence": 0.8472465253065024
282
+ },
283
+ {
284
+ "lower": 0.9,
285
+ "upper": 1.0,
286
+ "count": 74,
287
+ "accuracy": 0.9864864864864865,
288
+ "confidence": 0.977842163032285
289
+ }
290
+ ]
291
+ },
292
+ "held_out": {
293
+ "questions": 115,
294
+ "errors": 0,
295
+ "scenarios": 115,
296
+ "micro_accuracy": 0.8869565217391304,
297
+ "state_macro_accuracy": 0.8869565217391304,
298
+ "ece_10": 0.06774102719327177,
299
+ "confidence_gap": -0.017149033562759763,
300
+ "mean_confidence": 0.8698074881763707,
301
+ "latency_ms": {
302
+ "median": 1331.1896250088466,
303
+ "p95": 31440.3575410106
304
+ },
305
+ "option_count": {
306
+ "min": 2,
307
+ "max": 6,
308
+ "mean": 3.6782608695652175
309
+ },
310
+ "calibration_bins": [
311
+ {
312
+ "lower": 0.3,
313
+ "upper": 0.4,
314
+ "count": 5,
315
+ "accuracy": 0.4,
316
+ "confidence": 0.36867063350766766
317
+ },
318
+ {
319
+ "lower": 0.4,
320
+ "upper": 0.5,
321
+ "count": 3,
322
+ "accuracy": 0.3333333333333333,
323
+ "confidence": 0.4395554126130639
324
+ },
325
+ {
326
+ "lower": 0.5,
327
+ "upper": 0.6,
328
+ "count": 7,
329
+ "accuracy": 1.0,
330
+ "confidence": 0.5359182660921382
331
+ },
332
+ {
333
+ "lower": 0.6,
334
+ "upper": 0.7,
335
+ "count": 6,
336
+ "accuracy": 0.8333333333333334,
337
+ "confidence": 0.663160674101509
338
+ },
339
+ {
340
+ "lower": 0.7,
341
+ "upper": 0.8,
342
+ "count": 8,
343
+ "accuracy": 0.625,
344
+ "confidence": 0.7679741044507102
345
+ },
346
+ {
347
+ "lower": 0.8,
348
+ "upper": 0.9,
349
+ "count": 11,
350
+ "accuracy": 0.9090909090909091,
351
+ "confidence": 0.8677342210668939
352
+ },
353
+ {
354
+ "lower": 0.9,
355
+ "upper": 1.0,
356
+ "count": 75,
357
+ "accuracy": 0.96,
358
+ "confidence": 0.9792877408041276
359
+ }
360
+ ]
361
+ },
362
+ "all": {
363
+ "questions": 231,
364
+ "errors": 0,
365
+ "scenarios": 231,
366
+ "micro_accuracy": 0.8787878787878788,
367
+ "state_macro_accuracy": 0.8787878787878788,
368
+ "ece_10": 0.030691873313086204,
369
+ "confidence_gap": -0.008421448411867094,
370
+ "mean_confidence": 0.8703664303760117,
371
+ "latency_ms": {
372
+ "median": 1331.1896250088466,
373
+ "p95": 31386.2294999999
374
+ },
375
+ "option_count": {
376
+ "min": 2,
377
+ "max": 6,
378
+ "mean": 3.675324675324675
379
+ },
380
+ "calibration_bins": [
381
+ {
382
+ "lower": 0.3,
383
+ "upper": 0.4,
384
+ "count": 8,
385
+ "accuracy": 0.375,
386
+ "confidence": 0.3599872941163156
387
+ },
388
+ {
389
+ "lower": 0.4,
390
+ "upper": 0.5,
391
+ "count": 7,
392
+ "accuracy": 0.42857142857142855,
393
+ "confidence": 0.44999248257683433
394
+ },
395
+ {
396
+ "lower": 0.5,
397
+ "upper": 0.6,
398
+ "count": 12,
399
+ "accuracy": 0.8333333333333334,
400
+ "confidence": 0.5487722703386245
401
+ },
402
+ {
403
+ "lower": 0.6,
404
+ "upper": 0.7,
405
+ "count": 14,
406
+ "accuracy": 0.7142857142857143,
407
+ "confidence": 0.6646579291685718
408
+ },
409
+ {
410
+ "lower": 0.7,
411
+ "upper": 0.8,
412
+ "count": 18,
413
+ "accuracy": 0.6666666666666666,
414
+ "confidence": 0.7564103374085732
415
+ },
416
+ {
417
+ "lower": 0.8,
418
+ "upper": 0.9,
419
+ "count": 23,
420
+ "accuracy": 0.8695652173913043,
421
+ "confidence": 0.8570449884962548
422
+ },
423
+ {
424
+ "lower": 0.9,
425
+ "upper": 1.0,
426
+ "count": 149,
427
+ "accuracy": 0.9731543624161074,
428
+ "confidence": 0.9785698028503265
429
+ }
430
+ ]
431
+ }
432
+ }
433
+ }
calibration.json ADDED
@@ -0,0 +1,9 @@
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "method": "temperature_scaling",
3
+ "temperature": 1.11516790625,
4
+ "benchmark": "JevBench v1.2 public split (easy + original + hard)",
5
+ "fit_split": "even source rows; odd source rows held out for validation",
6
+ "fit_questions": 116,
7
+ "held_out_questions": 115,
8
+ "model": "Ruiruiz30/Jev-Omni-MLX-4bit"
9
+ }
omni_mlx/classifier.py CHANGED
@@ -24,8 +24,21 @@ def head_probabilities(hidden, head, count):
24
  return mx.softmax(logits, axis=-1)
25
 
26
 
 
 
 
 
 
 
 
 
 
 
 
 
 
27
  class Classifier:
28
- def __init__(self, path=".models/Jev-Omni-MLX-4bit"):
29
  self.path = Path(path).resolve()
30
  mx.set_cache_limit(256 * 1024**2)
31
  self.model = load_model(self.path, strict=True)
@@ -33,10 +46,20 @@ class Classifier:
33
  self.head = mx.load(str(self.path / "decision_head/weights.safetensors"))
34
  mx.eval(self.head)
35
  self.provenance = json.loads((self.path / "conversion.json").read_text())
 
 
 
 
 
 
 
 
 
 
36
 
37
  def predict(self, state, question, options, image=None, image_tokens=280):
38
- if not 2 <= len(options) <= 20 or len(set(options)) != len(options):
39
- raise ValueError("Supply 2–20 distinct options")
40
  if image_tokens not in (10, 20, 35, 70, 140, 280):
41
  raise ValueError("image_tokens must be 10, 20, 35, 70, 140, or 280")
42
  started = time.perf_counter()
@@ -55,8 +78,9 @@ class Classifier:
55
  prompts=formatted, add_special_tokens=False)
56
  inputs = {k: v.astype(mx.bfloat16) if isinstance(v, mx.array) and mx.issubdtype(v.dtype, mx.floating) else v for k,v in inputs.items()}
57
  ids = inputs["input_ids"]
58
- if ids.shape[1] > 1024:
59
- raise ValueError("Local prototype limited to 1024 input tokens to bound memory")
 
60
  # Batch size 1, no padding. Let the decoder build its causal/sliding and vision masks.
61
  extra = {k:v for k,v in inputs.items() if k not in {"input_ids", "attention_mask"}}
62
  mx.eval(inputs)
@@ -77,22 +101,99 @@ class Classifier:
77
  values = probabilities.tolist()
78
  if not all(0 <= value <= 1 for value in values):
79
  raise RuntimeError("Non-finite classifier output")
 
 
80
  elapsed_ms = (time.perf_counter() - started) * 1000
81
  best = max(range(len(values)), key=values.__getitem__)
82
  return {
83
  "model": "akhilaaa3/Jev-Omni", "backend": "mlx", "quantization_bits": self.provenance["bits"],
84
  "prediction": options[best], "prediction_index": best,
85
- "probabilities": dict(zip(options, values)), "calibrated": False,
 
86
  "metrics": {"elapsed_ms": elapsed_ms, "preprocessing_ms": preprocessing_ms,
87
  "vision_ms": vision_ms, "decoder_ms": decoder_ms,
88
  "input_tokens": ids.shape[1], "image_token_budget": image_tokens if image is not None else 0,
89
  "peak_metal_memory_gb": mx.get_peak_memory()/1e9, "generated_tokens": 0},
90
  }
91
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
92
 
93
  def main():
94
  parser = argparse.ArgumentParser(description=__doc__)
95
  parser.add_argument("--model", default=".models/Jev-Omni-MLX-4bit")
 
96
  parser.add_argument("--image")
97
  parser.add_argument("--state", default="")
98
  parser.add_argument("--question", required=True)
@@ -104,7 +205,7 @@ def main():
104
  if args.repeat < 1:
105
  parser.error("repeat must be positive")
106
  start = time.perf_counter()
107
- classifier = Classifier(args.model)
108
  load_ms = (time.perf_counter() - start) * 1000
109
  results = [classifier.predict(args.state, args.question, args.options, args.image, args.image_tokens) for _ in range(args.repeat)]
110
  result = {"load_ms": load_ms, "requests": results}
 
24
  return mx.softmax(logits, axis=-1)
25
 
26
 
27
+ def temperature_scale(values, temperature):
28
+ """Apply post-hoc temperature scaling to a probability vector."""
29
+ if temperature <= 0:
30
+ raise ValueError("calibration temperature must be positive")
31
+ import math
32
+ logits = [math.log(max(float(value), 1e-12)) for value in values]
33
+ scaled = [value / temperature for value in logits]
34
+ peak = max(scaled)
35
+ weights = [math.exp(value - peak) for value in scaled]
36
+ total = sum(weights)
37
+ return [value / total for value in weights]
38
+
39
+
40
  class Classifier:
41
+ def __init__(self, path=".models/Jev-Omni-MLX-4bit", calibration=None):
42
  self.path = Path(path).resolve()
43
  mx.set_cache_limit(256 * 1024**2)
44
  self.model = load_model(self.path, strict=True)
 
46
  self.head = mx.load(str(self.path / "decision_head/weights.safetensors"))
47
  mx.eval(self.head)
48
  self.provenance = json.loads((self.path / "conversion.json").read_text())
49
+ self.calibration = None
50
+ if calibration:
51
+ calibration_path = Path(calibration)
52
+ if not calibration_path.is_absolute():
53
+ calibration_path = self.path / calibration_path
54
+ self.calibration = json.loads(calibration_path.read_text())
55
+ if self.calibration.get("method") != "temperature_scaling":
56
+ raise ValueError("unsupported calibration method")
57
+ if float(self.calibration.get("temperature", 0)) <= 0:
58
+ raise ValueError("calibration temperature must be positive")
59
 
60
  def predict(self, state, question, options, image=None, image_tokens=280):
61
+ if not 2 <= len(options) <= 256 or len(set(options)) != len(options):
62
+ raise ValueError("Supply 2–256 distinct options")
63
  if image_tokens not in (10, 20, 35, 70, 140, 280):
64
  raise ValueError("image_tokens must be 10, 20, 35, 70, 140, or 280")
65
  started = time.perf_counter()
 
78
  prompts=formatted, add_special_tokens=False)
79
  inputs = {k: v.astype(mx.bfloat16) if isinstance(v, mx.array) and mx.issubdtype(v.dtype, mx.floating) else v for k,v in inputs.items()}
80
  ids = inputs["input_ids"]
81
+ # The upstream benchmark contains long state records. Keep the model's
82
+ # sliding/full-attention behavior intact instead of silently truncating
83
+ # benchmark inputs; the 16GB release has been tested up to 8k tokens.
84
  # Batch size 1, no padding. Let the decoder build its causal/sliding and vision masks.
85
  extra = {k:v for k,v in inputs.items() if k not in {"input_ids", "attention_mask"}}
86
  mx.eval(inputs)
 
101
  values = probabilities.tolist()
102
  if not all(0 <= value <= 1 for value in values):
103
  raise RuntimeError("Non-finite classifier output")
104
+ if self.calibration:
105
+ values = temperature_scale(values, float(self.calibration["temperature"]))
106
  elapsed_ms = (time.perf_counter() - started) * 1000
107
  best = max(range(len(values)), key=values.__getitem__)
108
  return {
109
  "model": "akhilaaa3/Jev-Omni", "backend": "mlx", "quantization_bits": self.provenance["bits"],
110
  "prediction": options[best], "prediction_index": best,
111
+ "probabilities": dict(zip(options, values)), "calibrated": bool(self.calibration),
112
+ "calibration": self.calibration if self.calibration else None,
113
  "metrics": {"elapsed_ms": elapsed_ms, "preprocessing_ms": preprocessing_ms,
114
  "vision_ms": vision_ms, "decoder_ms": decoder_ms,
115
  "input_tokens": ids.shape[1], "image_token_budget": image_tokens if image is not None else 0,
116
  "peak_metal_memory_gb": mx.get_peak_memory()/1e9, "generated_tokens": 0},
117
  }
118
 
119
+ def predict_many_text(self, state, requests):
120
+ """Score several text questions sharing one state with a prefix KV cache.
121
+
122
+ DecisionBench puts multiple typed questions on each long state. Reusing
123
+ the common prefix preserves the logits while avoiding repeated prefill.
124
+ This path is text-only; image requests continue through ``predict``.
125
+ """
126
+ if not requests:
127
+ return []
128
+ if any(not 2 <= len(options) <= 256 or len(set(options)) != len(options)
129
+ for _, options in requests):
130
+ raise ValueError("Supply 2–256 distinct options")
131
+ started = time.perf_counter()
132
+ formatted_prompts = [self.processor.apply_chat_template(
133
+ [{"role": "user", "content": [{"type": "text", "text": prompt(state, question, options)}]}],
134
+ add_generation_prompt=True, tokenize=False, enable_thinking=False,
135
+ ) for question, options in requests]
136
+ prepared = []
137
+ for formatted in formatted_prompts:
138
+ inputs = prepare_inputs(self.processor, images=None, prompts=formatted,
139
+ add_special_tokens=False)
140
+ inputs = {k: v.astype(mx.bfloat16) if isinstance(v, mx.array) and mx.issubdtype(v.dtype, mx.floating) else v
141
+ for k, v in inputs.items()}
142
+ embedded = self.model.get_input_embeddings(
143
+ input_ids=inputs["input_ids"],
144
+ **{k: v for k, v in inputs.items() if k not in {"input_ids", "attention_mask"}},
145
+ )
146
+ prepared.append((inputs, embedded))
147
+ mx.eval([item for inputs, embedded in prepared for item in (embedded.inputs_embeds, embedded.per_layer_inputs)
148
+ if item is not None])
149
+ ids = [inputs["input_ids"][0].tolist() for inputs, _ in prepared]
150
+ prefix_len = min(len(ids[0]), *(len(row) for row in ids))
151
+ while prefix_len and any(row[:prefix_len] != ids[0][:prefix_len] for row in ids[1:]):
152
+ prefix_len -= 1
153
+ if prefix_len == 0:
154
+ return [self.predict(state, question, options) for question, options in requests]
155
+ prefix_cache = self.model.make_cache()
156
+ first_inputs, first_embedded = prepared[0]
157
+ prefix_kwargs = {"inputs_embeds": first_embedded.inputs_embeds[:, :prefix_len],
158
+ "cache": prefix_cache}
159
+ if first_embedded.per_layer_inputs is not None:
160
+ prefix_kwargs["per_layer_inputs"] = first_embedded.per_layer_inputs[:, :prefix_len]
161
+ prefix_hidden = self.model.language_model.model(**prefix_kwargs)
162
+ mx.eval(prefix_hidden)
163
+ snapshots = [cache.prefix_cache_snapshot() for cache in prefix_cache]
164
+ prefix_ms = (time.perf_counter() - started) * 1000
165
+ output = []
166
+ for (question, options), (inputs, embedded) in zip(requests, prepared):
167
+ caches = self.model.make_cache()
168
+ for cache, snapshot in zip(caches, snapshots):
169
+ cache.prefix_cache_restore(snapshot)
170
+ suffix_kwargs = {"inputs_embeds": embedded.inputs_embeds[:, prefix_len:], "cache": caches}
171
+ if embedded.per_layer_inputs is not None:
172
+ suffix_kwargs["per_layer_inputs"] = embedded.per_layer_inputs[:, prefix_len:]
173
+ decoder_started = time.perf_counter()
174
+ hidden = self.model.language_model.model(**suffix_kwargs)
175
+ probabilities = head_probabilities(hidden[:, -1], self.head, len(options))[0]
176
+ mx.eval(probabilities)
177
+ values = probabilities.tolist()
178
+ best = max(range(len(values)), key=values.__getitem__)
179
+ output.append({
180
+ "model": "akhilaaa3/Jev-Omni", "backend": "mlx", "quantization_bits": self.provenance["bits"],
181
+ "prediction": options[best], "prediction_index": best,
182
+ "probabilities": dict(zip(options, values)), "calibrated": False, "calibration": None,
183
+ "metrics": {"elapsed_ms": (time.perf_counter() - decoder_started) * 1000,
184
+ "preprocessing_ms": prefix_ms / len(requests), "vision_ms": 0,
185
+ "decoder_ms": (time.perf_counter() - decoder_started) * 1000,
186
+ "input_tokens": inputs["input_ids"].shape[1], "image_token_budget": 0,
187
+ "peak_metal_memory_gb": mx.get_peak_memory() / 1e9, "generated_tokens": 0,
188
+ "prefix_cache_reused": True, "prefix_tokens": prefix_len},
189
+ })
190
+ return output
191
+
192
 
193
  def main():
194
  parser = argparse.ArgumentParser(description=__doc__)
195
  parser.add_argument("--model", default=".models/Jev-Omni-MLX-4bit")
196
+ parser.add_argument("--calibration", help="JSON temperature-scaling calibration file")
197
  parser.add_argument("--image")
198
  parser.add_argument("--state", default="")
199
  parser.add_argument("--question", required=True)
 
205
  if args.repeat < 1:
206
  parser.error("repeat must be positive")
207
  start = time.perf_counter()
208
+ classifier = Classifier(args.model, args.calibration)
209
  load_ms = (time.perf_counter() - start) * 1000
210
  results = [classifier.predict(args.state, args.question, args.options, args.image, args.image_tokens) for _ in range(args.repeat)]
211
  result = {"load_ms": load_ms, "requests": results}