Mapika commited on
Commit
4535722
·
verified ·
1 Parent(s): 51c6fbc

Upload folder using huggingface_hub

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
+ tokenizer.json filter=lfs diff=lfs merge=lfs -text
README.md ADDED
@@ -0,0 +1,240 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: apache-2.0
3
+ base_model: Qwen/Qwen3.5-2B-Base
4
+ language: [en]
5
+ pipeline_tag: text-classification
6
+ tags: [decision-model, calibrated, structured-output, multi-task, system-one, one-pass]
7
+ ---
8
+
9
+ # decider-decider-2B: typed decisions with calibrated probabilities in one forward pass
10
+
11
+ An open replication of the "System One model" idea: a language model that does not
12
+ generate text. It reads a context plus one or more typed questions, each with an
13
+ explicit option list, and returns a probability distribution over the options for
14
+ every question from a single forward pass. No decoding, no JSON parsing, no
15
+ schema violations. It is meant to be called from software, not chatted with.
16
+
17
+ Base model: [Qwen/Qwen3.5-2B-Base](https://huggingface.co/Qwen/Qwen3.5-2B-Base) (1.9B parameters),
18
+ fully fine-tuned for one epoch (942k examples, 183M tokens, 2.5 hours on one
19
+ NVIDIA GH200) with cross-entropy, a proper scoring rule, on a mixture of 64
20
+ public decision datasets.
21
+
22
+ ## Usage
23
+
24
+ ```python
25
+ from decider.infer import Decider # decider/ is included in this repo
26
+ d = Decider("<this repo>")
27
+ d.decide("My card was charged twice for the same purchase.",
28
+ [{"question": "Which department should handle this?", "options": ["billing", "technical support", "sales"]},
29
+ {"question": "Does this need a refund action?", "options": ["no", "yes"]}])
30
+ # [{'choice': 'billing', 'confidence': 0.99, 'probs': {...}}, {'choice': 'yes', 'confidence': 0.99, 'probs': {...}}]
31
+ ```
32
+
33
+ `decide_batch` scores many contexts, each with many questions, in one call.
34
+ Set `abstain_below=t` to return `None` for decisions with confidence under `t`
35
+ (route to a human). 2 to 10 options per question.
36
+
37
+ Requirements: `torch`, `transformers>=5`, and `flash-linear-attention` (Triton
38
+ kernels for the Qwen3.5 linear-attention layers; the model runs without it but
39
+ several times slower). Python 3.11+ recommended so those kernels can use
40
+ `torch.compile`.
41
+
42
+ Without the helper package, the same computation in plain `transformers`:
43
+
44
+ ```python
45
+ import torch
46
+ from transformers import AutoTokenizer, AutoModelForCausalLM
47
+ tok = AutoTokenizer.from_pretrained(REPO); m = AutoModelForCausalLM.from_pretrained(REPO, dtype=torch.bfloat16).cuda().eval()
48
+ prompt = ("Context:\nMy card was charged twice for the same purchase.\n\n"
49
+ "Question: Which department should handle this?\nOptions:\n(A) billing\n(B) technical support\n(C) sales\nAnswer: (")
50
+ ids = tok(prompt, return_tensors="pt").to("cuda")
51
+ with torch.no_grad():
52
+ logits = m(**ids).logits[0, -1]
53
+ letters = [tok.encode(L, add_special_tokens=False)[0] for L in "ABC"]
54
+ probs = torch.softmax(logits[letters].float(), -1) # -> P(billing), P(technical support), P(sales)
55
+ ```
56
+
57
+ For several questions in one pass, append further `Question k: ... Answer k: (`
58
+ blocks and read the logits at each `(` position (see `decider/prompt.py`).
59
+
60
+ ## How it works
61
+
62
+ Prompt: `Context: ...` followed by, for each question, the question text, the
63
+ numbered options `(A) ... (B) ...`, and an answer slot `Answer k: (`. The hidden
64
+ state at each slot is projected with the option-letter rows of the LM head and
65
+ softmaxed over the valid letters. Letters are never generated, so all slots are
66
+ read from one pass. Large label sets were sub-sampled to at most 10 options per
67
+ training example (gold always kept, order shuffled), so the model conditions on
68
+ the supplied candidates rather than a fixed head.
69
+
70
+ ## Field types
71
+
72
+ * **bool** (`noul`): probability of "yes".
73
+ * **choice**: argmax option, its probability, and the full distribution.
74
+ * **scale**: an ordered legend (e.g. 0: none ... 3: high); returns the expected
75
+ level (`score`), the probability of the most likely level, and the distribution.
76
+
77
+ ## Training data
78
+
79
+ 64 public datasets, up to 20k examples each (`decider/data.py`, `decider/data2.py`):
80
+ intent detection, ticket routing, topic classification, sentiment, emotion,
81
+ moderation (toxicity, hate, spam, jailbreak, safety), NLI, paraphrase, fact
82
+ verification, passage relevance, reading comprehension, multiple-choice QA,
83
+ ordinal rating scales (HelpSteer2 attributes, STS-B, hate-speech intensity,
84
+ LIAR2 truthfulness), pairwise response preference (HelpSteer3, UltraFeedback,
85
+ SHP, HH-RLHF) and tool selection (Glaive, ToolACE).
86
+ Abstention augmentation: in 10% of questions with three or more options the
87
+ gold option is removed and "none of the above" becomes the answer.
88
+
89
+ ## Evaluation
90
+
91
+ | Model | Split | Acc | NLL | Brier | ECE | AURC | Acc@80% |
92
+ |---|---|---|---|---|---|---|---|
93
+ | Qwen3.5-2B-Base, zero-shot | in-task (64) | 0.620 | 0.908 | 0.493 | 0.121 | 0.280 | 0.663 |
94
+ | Qwen3.5-2B-Base, zero-shot | held-out (23) | 0.642 | 0.853 | 0.460 | 0.105 | 0.242 | 0.685 |
95
+ | Qwen3.5-4B-Base, zero-shot | in-task (64) | 0.695 | 0.768 | 0.405 | 0.090 | 0.206 | 0.742 |
96
+ | Qwen3.5-4B-Base, zero-shot | held-out (23) | 0.711 | 0.734 | 0.390 | 0.089 | 0.169 | 0.761 |
97
+ | **this model** | in-task (64) | 0.807 | 0.460 | 0.257 | 0.029 | 0.097 | 0.859 |
98
+ | **this model** | held-out (23) | 0.745 | 0.634 | 0.343 | 0.075 | 0.138 | 0.802 |
99
+
100
+
101
+ Per-task accuracy / ECE on the held-out datasets:
102
+
103
+ | Task | Qwen3.5-2B-Base, zero-shot | Qwen3.5-4B-Base, zero-shot | this model |
104
+ |---|---|---|---|
105
+ | abstain_probe | 0.377 / 0.193 | 0.453 / 0.242 | 0.785 / 0.037 |
106
+ | ade | 0.794 / 0.117 | 0.746 / 0.101 | 0.808 / 0.038 |
107
+ | arena_pref | 0.382 / 0.228 | 0.434 / 0.159 | 0.474 / 0.144 |
108
+ | bbc_news | 0.910 / 0.010 | 0.927 / 0.024 | 0.941 / 0.016 |
109
+ | cb | 0.554 / 0.115 | 0.696 / 0.087 | 0.893 / 0.115 |
110
+ | cr_reviews | 0.882 / 0.073 | 0.923 / 0.015 | 0.902 / 0.020 |
111
+ | dolly_category | 0.269 / 0.137 | 0.351 / 0.149 | 0.318 / 0.211 |
112
+ | fin_phrasebank | 0.560 / 0.032 | 0.713 / 0.050 | 0.660 / 0.082 |
113
+ | fin_sentiment | 0.345 / 0.335 | 0.437 / 0.303 | 0.769 / 0.031 |
114
+ | hermes_tools | 0.704 / 0.067 | 0.702 / 0.160 | 0.737 / 0.155 |
115
+ | massive_scenario | 0.706 / 0.094 | 0.707 / 0.032 | 0.815 / 0.023 |
116
+ | paws | 0.759 / 0.123 | 0.834 / 0.018 | 0.691 / 0.221 |
117
+ | pubmedqa | 0.728 / 0.042 | 0.768 / 0.082 | 0.790 / 0.063 |
118
+ | quality | 0.457 / 0.205 | 0.519 / 0.184 | 0.513 / 0.137 |
119
+ | reward_bench | 0.624 / 0.074 | 0.772 / 0.045 | 0.799 / 0.043 |
120
+ | sciq | 0.978 / 0.011 | 0.988 / 0.019 | 0.982 / 0.017 |
121
+ | social_iqa | 0.659 / 0.089 | 0.739 / 0.053 | 0.712 / 0.060 |
122
+ | strategyqa | 0.566 / 0.029 | 0.646 / 0.037 | 0.613 / 0.069 |
123
+ | student_questions | 0.903 / 0.059 | 0.933 / 0.021 | 0.927 / 0.030 |
124
+ | trec | 0.696 / 0.062 | 0.822 / 0.050 | 0.766 / 0.036 |
125
+ | truthfulqa | 0.460 / 0.111 | 0.591 / 0.113 | 0.497 / 0.091 |
126
+ | tweet_irony | 0.511 / 0.158 | 0.676 / 0.073 | 0.779 / 0.059 |
127
+ | xstory_cloze | 0.941 / 0.061 | 0.981 / 0.041 | 0.965 / 0.028 |
128
+
129
+ | Task | Qwen3.5-2B-Base, zero-shot | Qwen3.5-4B-Base, zero-shot | this model |
130
+ |---|---|---|---|
131
+ | abstain_probe | 0.377 / 0.193 | 0.453 / 0.242 | 0.785 / 0.037 |
132
+ | ade | 0.794 / 0.117 | 0.746 / 0.101 | 0.808 / 0.038 |
133
+ | arena_pref | 0.382 / 0.228 | 0.434 / 0.159 | 0.474 / 0.144 |
134
+ | bbc_news | 0.910 / 0.010 | 0.927 / 0.024 | 0.941 / 0.016 |
135
+ | cb | 0.554 / 0.115 | 0.696 / 0.087 | 0.893 / 0.115 |
136
+ | cr_reviews | 0.882 / 0.073 | 0.923 / 0.015 | 0.902 / 0.020 |
137
+ | dolly_category | 0.269 / 0.137 | 0.351 / 0.149 | 0.318 / 0.211 |
138
+ | fin_phrasebank | 0.560 / 0.032 | 0.713 / 0.050 | 0.660 / 0.082 |
139
+ | fin_sentiment | 0.345 / 0.335 | 0.437 / 0.303 | 0.769 / 0.031 |
140
+ | hermes_tools | 0.704 / 0.067 | 0.702 / 0.160 | 0.737 / 0.155 |
141
+ | massive_scenario | 0.706 / 0.094 | 0.707 / 0.032 | 0.815 / 0.023 |
142
+ | paws | 0.759 / 0.123 | 0.834 / 0.018 | 0.691 / 0.221 |
143
+ | pubmedqa | 0.728 / 0.042 | 0.768 / 0.082 | 0.790 / 0.063 |
144
+ | quality | 0.457 / 0.205 | 0.519 / 0.184 | 0.513 / 0.137 |
145
+ | reward_bench | 0.624 / 0.074 | 0.772 / 0.045 | 0.799 / 0.043 |
146
+ | sciq | 0.978 / 0.011 | 0.988 / 0.019 | 0.982 / 0.017 |
147
+ | social_iqa | 0.659 / 0.089 | 0.739 / 0.053 | 0.712 / 0.060 |
148
+ | strategyqa | 0.566 / 0.029 | 0.646 / 0.037 | 0.613 / 0.069 |
149
+ | student_questions | 0.903 / 0.059 | 0.933 / 0.021 | 0.927 / 0.030 |
150
+ | trec | 0.696 / 0.062 | 0.822 / 0.050 | 0.766 / 0.036 |
151
+ | truthfulqa | 0.460 / 0.111 | 0.591 / 0.113 | 0.497 / 0.091 |
152
+ | tweet_irony | 0.511 / 0.158 | 0.676 / 0.073 | 0.779 / 0.059 |
153
+ | xstory_cloze | 0.941 / 0.061 | 0.981 / 0.041 | 0.965 / 0.028 |
154
+
155
+ | Task | Qwen3.5-2B-Base, zero-shot | Qwen3.5-4B-Base, zero-shot | this model |
156
+ |---|---|---|---|
157
+ | ade | 0.794 / 0.117 | 0.746 / 0.101 | 0.827 / 0.023 |
158
+ | bbc_news | 0.910 / 0.010 | 0.927 / 0.024 | 0.919 / 0.025 |
159
+ | cr_reviews | 0.882 / 0.073 | 0.923 / 0.015 | 0.911 / 0.019 |
160
+ | dolly_category | 0.269 / 0.137 | 0.351 / 0.149 | 0.316 / 0.212 |
161
+ | fin_phrasebank | 0.560 / 0.032 | 0.713 / 0.050 | 0.640 / 0.125 |
162
+ | fin_sentiment | 0.345 / 0.335 | 0.437 / 0.303 | 0.773 / 0.029 |
163
+ | massive_scenario | 0.706 / 0.094 | 0.707 / 0.032 | 0.822 / 0.017 |
164
+ | paws | 0.759 / 0.123 | 0.834 / 0.018 | 0.693 / 0.189 |
165
+ | pubmedqa | 0.728 / 0.042 | 0.768 / 0.082 | 0.768 / 0.043 |
166
+ | sciq | 0.978 / 0.011 | 0.988 / 0.019 | 0.985 / 0.016 |
167
+ | social_iqa | 0.659 / 0.089 | 0.739 / 0.053 | 0.712 / 0.055 |
168
+ | strategyqa | 0.566 / 0.029 | 0.646 / 0.037 | 0.594 / 0.095 |
169
+ | student_questions | 0.903 / 0.059 | 0.933 / 0.021 | 0.930 / 0.017 |
170
+ | trec | 0.696 / 0.062 | 0.822 / 0.050 | 0.772 / 0.044 |
171
+ | truthfulqa | 0.460 / 0.111 | 0.591 / 0.113 | 0.541 / 0.062 |
172
+ | tweet_irony | 0.511 / 0.158 | 0.676 / 0.073 | 0.754 / 0.054 |
173
+
174
+ | Task | Qwen3.5-2B-Base, zero-shot | Qwen3.5-4B-Base, zero-shot | this model (200k-example run) |
175
+ |---|---|---|---|
176
+ | ade | 0.794 / 0.117 | 0.746 / 0.101 | 0.808 / 0.046 |
177
+ | bbc_news | 0.910 / 0.010 | 0.927 / 0.024 | 0.928 / 0.014 |
178
+ | cr_reviews | 0.882 / 0.073 | 0.923 / 0.015 | 0.915 / 0.012 |
179
+ | dolly_category | 0.269 / 0.137 | 0.351 / 0.149 | 0.325 / 0.193 |
180
+ | fin_phrasebank | 0.560 / 0.032 | 0.713 / 0.050 | 0.652 / 0.108 |
181
+ | fin_sentiment | 0.345 / 0.335 | 0.437 / 0.303 | 0.796 / 0.042 |
182
+ | massive_scenario | 0.706 / 0.094 | 0.707 / 0.032 | 0.808 / 0.022 |
183
+ | paws | 0.759 / 0.123 | 0.834 / 0.018 | 0.713 / 0.139 |
184
+ | pubmedqa | 0.728 / 0.042 | 0.768 / 0.082 | 0.772 / 0.048 |
185
+ | sciq | 0.978 / 0.011 | 0.988 / 0.019 | 0.983 / 0.021 |
186
+ | social_iqa | 0.659 / 0.089 | 0.739 / 0.053 | 0.703 / 0.060 |
187
+ | strategyqa | 0.566 / 0.029 | 0.646 / 0.037 | 0.597 / 0.089 |
188
+ | student_questions | 0.903 / 0.059 | 0.933 / 0.021 | 0.927 / 0.029 |
189
+ | trec | 0.696 / 0.062 | 0.822 / 0.050 | 0.816 / 0.039 |
190
+ | truthfulqa | 0.460 / 0.111 | 0.591 / 0.113 | 0.529 / 0.058 |
191
+ | tweet_irony | 0.511 / 0.158 | 0.676 / 0.073 | 0.769 / 0.048 |
192
+
193
+
194
+ *In-task* = test splits of the 64 training datasets. *Held-out* = 23 datasets never
195
+ seen in training: TREC, BBC news, PAWS, SciQ, Social IQa, StrategyQA, PubMedQA,
196
+ TruthfulQA, tweet irony, financial sentiment, ADE, MASSIVE scenario, student
197
+ question categories, Dolly categories, CR reviews, Financial PhraseBank,
198
+ CommitmentBank, QuALITY, XStoryCloze, RewardBench, Arena preferences (3-way),
199
+ Hermes tool selection, and an abstention probe (held-out classification tasks
200
+ where in half the cases the correct option is absent and "none of the above" is
201
+ right). Chance accuracy is 0.33 on both sets. ECE = expected calibration error
202
+ (15 bins), AURC = area under the risk-coverage curve, acc@80 = accuracy on
203
+ the 80% most confident decisions.
204
+
205
+ ## Speed
206
+
207
+ One NVIDIA GH200, bf16. `decider.infer.Decider` uses shape-bucketed CUDA
208
+ graphs (`decider/engine.py`); the micro-batching server is `decider/serve.py`
209
+ in the GitHub repo. Support-ticket contexts of ~230 tokens with 3 to 5 typed
210
+ questions each:
211
+
212
+ | setting | p50 latency | throughput |
213
+ |---|---|---|
214
+ | single request, eager PyTorch | 49 ms | |
215
+ | single request, CUDA-graph engine | 6.6 ms | |
216
+ | batch of 32, in-process | 115 ms | ~840 decisions/s |
217
+ | HTTP server, 1 client | 10 ms | 97 req/s |
218
+ | HTTP server, 64 clients | 231 ms | 263 req/s, 1314 decisions/s |
219
+
220
+ ## Limitations
221
+
222
+ * English only. Options must be short phrases; free-text fields are not supported.
223
+ * Calibration is measured on public datasets; verify it on your own labelled
224
+ data before using confidence for routing.
225
+ * No reasoning: this is a fast pattern-matching decision model, not a chat model.
226
+ * Maximum context 1536 tokens as trained.
227
+ * Knowledge-heavy multiple choice (MMLU, MedQA, ARC) improves only modestly over the
228
+ base model; fine-tuning on decisions does not add world knowledge.
229
+ * Compared with a run on the 47-dataset v1 mixture, adding the v2 datasets
230
+ raised held-out accuracy but lowered two held-out tasks: Hermes tool selection
231
+ (0.80 to 0.74) and TruthfulQA (0.54 to 0.50).
232
+ * Scale fields are the least trained type; expect wider distributions there.
233
+ * One in-task dataset, `tweet_hate` (SemEval-2019 HatEval), stays near chance on its
234
+ test split. That split is known to differ from its training split in collection
235
+ and label definition; the number is reported as measured.
236
+
237
+ ## Reproduction
238
+
239
+ Code, data registry, training and evaluation scripts: https://github.com/Mapika/decider
240
+ (`decider/` in this model repo is the inference subset of that package).
chat_template.jinja ADDED
@@ -0,0 +1,154 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {%- set image_count = namespace(value=0) %}
2
+ {%- set video_count = namespace(value=0) %}
3
+ {%- macro render_content(content, do_vision_count, is_system_content=false) %}
4
+ {%- if content is string %}
5
+ {{- content }}
6
+ {%- elif content is iterable and content is not mapping %}
7
+ {%- for item in content %}
8
+ {%- if 'image' in item or 'image_url' in item or item.type == 'image' %}
9
+ {%- if is_system_content %}
10
+ {{- raise_exception('System message cannot contain images.') }}
11
+ {%- endif %}
12
+ {%- if do_vision_count %}
13
+ {%- set image_count.value = image_count.value + 1 %}
14
+ {%- endif %}
15
+ {%- if add_vision_id %}
16
+ {{- 'Picture ' ~ image_count.value ~ ': ' }}
17
+ {%- endif %}
18
+ {{- '<|vision_start|><|image_pad|><|vision_end|>' }}
19
+ {%- elif 'video' in item or item.type == 'video' %}
20
+ {%- if is_system_content %}
21
+ {{- raise_exception('System message cannot contain videos.') }}
22
+ {%- endif %}
23
+ {%- if do_vision_count %}
24
+ {%- set video_count.value = video_count.value + 1 %}
25
+ {%- endif %}
26
+ {%- if add_vision_id %}
27
+ {{- 'Video ' ~ video_count.value ~ ': ' }}
28
+ {%- endif %}
29
+ {{- '<|vision_start|><|video_pad|><|vision_end|>' }}
30
+ {%- elif 'text' in item %}
31
+ {{- item.text }}
32
+ {%- else %}
33
+ {{- raise_exception('Unexpected item type in content.') }}
34
+ {%- endif %}
35
+ {%- endfor %}
36
+ {%- elif content is none or content is undefined %}
37
+ {{- '' }}
38
+ {%- else %}
39
+ {{- raise_exception('Unexpected content type.') }}
40
+ {%- endif %}
41
+ {%- endmacro %}
42
+ {%- if not messages %}
43
+ {{- raise_exception('No messages provided.') }}
44
+ {%- endif %}
45
+ {%- if tools and tools is iterable and tools is not mapping %}
46
+ {{- '<|im_start|>system\n' }}
47
+ {{- "# Tools\n\nYou have access to the following functions:\n\n<tools>" }}
48
+ {%- for tool in tools %}
49
+ {{- "\n" }}
50
+ {{- tool | tojson }}
51
+ {%- endfor %}
52
+ {{- "\n</tools>" }}
53
+ {{- '\n\nIf you choose to call a function ONLY reply in the following format with NO suffix:\n\n<tool_call>\n<function=example_function_name>\n<parameter=example_parameter_1>\nvalue_1\n</parameter>\n<parameter=example_parameter_2>\nThis is the value for the second parameter\nthat can span\nmultiple lines\n</parameter>\n</function>\n</tool_call>\n\n<IMPORTANT>\nReminder:\n- Function calls MUST follow the specified format: an inner <function=...></function> block must be nested within <tool_call></tool_call> XML tags\n- Required parameters MUST be specified\n- You may provide optional reasoning for your function call in natural language BEFORE the function call, but NOT after\n- If there is no function call available, answer the question like normal with your current knowledge and do not tell the user about function calls\n</IMPORTANT>' }}
54
+ {%- if messages[0].role == 'system' %}
55
+ {%- set content = render_content(messages[0].content, false, true)|trim %}
56
+ {%- if content %}
57
+ {{- '\n\n' + content }}
58
+ {%- endif %}
59
+ {%- endif %}
60
+ {{- '<|im_end|>\n' }}
61
+ {%- else %}
62
+ {%- if messages[0].role == 'system' %}
63
+ {%- set content = render_content(messages[0].content, false, true)|trim %}
64
+ {{- '<|im_start|>system\n' + content + '<|im_end|>\n' }}
65
+ {%- endif %}
66
+ {%- endif %}
67
+ {%- set ns = namespace(multi_step_tool=true, last_query_index=messages|length - 1) %}
68
+ {%- for message in messages[::-1] %}
69
+ {%- set index = (messages|length - 1) - loop.index0 %}
70
+ {%- if ns.multi_step_tool and message.role == "user" %}
71
+ {%- set content = render_content(message.content, false)|trim %}
72
+ {%- if not(content.startswith('<tool_response>') and content.endswith('</tool_response>')) %}
73
+ {%- set ns.multi_step_tool = false %}
74
+ {%- set ns.last_query_index = index %}
75
+ {%- endif %}
76
+ {%- endif %}
77
+ {%- endfor %}
78
+ {%- if ns.multi_step_tool %}
79
+ {{- raise_exception('No user query found in messages.') }}
80
+ {%- endif %}
81
+ {%- for message in messages %}
82
+ {%- set content = render_content(message.content, true)|trim %}
83
+ {%- if message.role == "system" %}
84
+ {%- if not loop.first %}
85
+ {{- raise_exception('System message must be at the beginning.') }}
86
+ {%- endif %}
87
+ {%- elif message.role == "user" %}
88
+ {{- '<|im_start|>' + message.role + '\n' + content + '<|im_end|>' + '\n' }}
89
+ {%- elif message.role == "assistant" %}
90
+ {%- set reasoning_content = '' %}
91
+ {%- if message.reasoning_content is string %}
92
+ {%- set reasoning_content = message.reasoning_content %}
93
+ {%- else %}
94
+ {%- if '</think>' in content %}
95
+ {%- set reasoning_content = content.split('</think>')[0].rstrip('\n').split('<think>')[-1].lstrip('\n') %}
96
+ {%- set content = content.split('</think>')[-1].lstrip('\n') %}
97
+ {%- endif %}
98
+ {%- endif %}
99
+ {%- set reasoning_content = reasoning_content|trim %}
100
+ {%- if loop.index0 > ns.last_query_index %}
101
+ {{- '<|im_start|>' + message.role + '\n<think>\n' + reasoning_content + '\n</think>\n\n' + content }}
102
+ {%- else %}
103
+ {{- '<|im_start|>' + message.role + '\n' + content }}
104
+ {%- endif %}
105
+ {%- if message.tool_calls and message.tool_calls is iterable and message.tool_calls is not mapping %}
106
+ {%- for tool_call in message.tool_calls %}
107
+ {%- if tool_call.function is defined %}
108
+ {%- set tool_call = tool_call.function %}
109
+ {%- endif %}
110
+ {%- if loop.first %}
111
+ {%- if content|trim %}
112
+ {{- '\n\n<tool_call>\n<function=' + tool_call.name + '>\n' }}
113
+ {%- else %}
114
+ {{- '<tool_call>\n<function=' + tool_call.name + '>\n' }}
115
+ {%- endif %}
116
+ {%- else %}
117
+ {{- '\n<tool_call>\n<function=' + tool_call.name + '>\n' }}
118
+ {%- endif %}
119
+ {%- if tool_call.arguments is defined %}
120
+ {%- for args_name, args_value in tool_call.arguments|items %}
121
+ {{- '<parameter=' + args_name + '>\n' }}
122
+ {%- set args_value = args_value | tojson | safe if args_value is mapping or (args_value is sequence and args_value is not string) else args_value | string %}
123
+ {{- args_value }}
124
+ {{- '\n</parameter>\n' }}
125
+ {%- endfor %}
126
+ {%- endif %}
127
+ {{- '</function>\n</tool_call>' }}
128
+ {%- endfor %}
129
+ {%- endif %}
130
+ {{- '<|im_end|>\n' }}
131
+ {%- elif message.role == "tool" %}
132
+ {%- if loop.previtem and loop.previtem.role != "tool" %}
133
+ {{- '<|im_start|>user' }}
134
+ {%- endif %}
135
+ {{- '\n<tool_response>\n' }}
136
+ {{- content }}
137
+ {{- '\n</tool_response>' }}
138
+ {%- if not loop.last and loop.nextitem.role != "tool" %}
139
+ {{- '<|im_end|>\n' }}
140
+ {%- elif loop.last %}
141
+ {{- '<|im_end|>\n' }}
142
+ {%- endif %}
143
+ {%- else %}
144
+ {{- raise_exception('Unexpected message role.') }}
145
+ {%- endif %}
146
+ {%- endfor %}
147
+ {%- if add_generation_prompt %}
148
+ {{- '<|im_start|>assistant\n' }}
149
+ {%- if enable_thinking is defined and enable_thinking is true %}
150
+ {{- '<think>\n' }}
151
+ {%- else %}
152
+ {{- '<think>\n\n</think>\n\n' }}
153
+ {%- endif %}
154
+ {%- endif %}
config.json ADDED
@@ -0,0 +1,75 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "architectures": [
3
+ "Qwen3_5ForCausalLM"
4
+ ],
5
+ "attention_bias": false,
6
+ "attention_dropout": 0.0,
7
+ "attn_output_gate": true,
8
+ "bos_token_id": null,
9
+ "dtype": "bfloat16",
10
+ "eos_token_id": 248044,
11
+ "full_attention_interval": 4,
12
+ "head_dim": 256,
13
+ "hidden_act": "silu",
14
+ "hidden_size": 2048,
15
+ "initializer_range": 0.02,
16
+ "intermediate_size": 6144,
17
+ "layer_types": [
18
+ "linear_attention",
19
+ "linear_attention",
20
+ "linear_attention",
21
+ "full_attention",
22
+ "linear_attention",
23
+ "linear_attention",
24
+ "linear_attention",
25
+ "full_attention",
26
+ "linear_attention",
27
+ "linear_attention",
28
+ "linear_attention",
29
+ "full_attention",
30
+ "linear_attention",
31
+ "linear_attention",
32
+ "linear_attention",
33
+ "full_attention",
34
+ "linear_attention",
35
+ "linear_attention",
36
+ "linear_attention",
37
+ "full_attention",
38
+ "linear_attention",
39
+ "linear_attention",
40
+ "linear_attention",
41
+ "full_attention"
42
+ ],
43
+ "linear_conv_kernel_dim": 4,
44
+ "linear_key_head_dim": 128,
45
+ "linear_num_key_heads": 16,
46
+ "linear_num_value_heads": 16,
47
+ "linear_value_head_dim": 128,
48
+ "mamba_ssm_dtype": "float32",
49
+ "max_position_embeddings": 262144,
50
+ "mlp_only_layers": [],
51
+ "model_type": "qwen3_5_text",
52
+ "mtp_num_hidden_layers": 1,
53
+ "mtp_use_dedicated_embeddings": false,
54
+ "num_attention_heads": 8,
55
+ "num_hidden_layers": 24,
56
+ "num_key_value_heads": 2,
57
+ "pad_token_id": null,
58
+ "partial_rotary_factor": 0.25,
59
+ "rms_norm_eps": 1e-06,
60
+ "rope_parameters": {
61
+ "mrope_interleaved": true,
62
+ "mrope_section": [
63
+ 11,
64
+ 11,
65
+ 10
66
+ ],
67
+ "partial_rotary_factor": 0.25,
68
+ "rope_theta": 10000000,
69
+ "rope_type": "default"
70
+ },
71
+ "tie_word_embeddings": true,
72
+ "transformers_version": "5.17.0",
73
+ "use_cache": true,
74
+ "vocab_size": 248320
75
+ }
decider/__init__.py ADDED
File without changes
decider/__pycache__/__init__.cpython-312.pyc ADDED
Binary file (174 Bytes). View file
 
decider/__pycache__/infer.cpython-312.pyc ADDED
Binary file (13.5 kB). View file
 
decider/__pycache__/model.cpython-312.pyc ADDED
Binary file (4.86 kB). View file
 
decider/__pycache__/prompt.cpython-312.pyc ADDED
Binary file (3.66 kB). View file
 
decider/engine.py ADDED
@@ -0,0 +1,124 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Low-latency inference engine: shape-bucketed CUDA graphs over the one-pass decision model.
2
+
3
+ Right padding + causal layers => pad positions never influence earlier slots, so no attention
4
+ mask is needed and every (B, T) bucket can be captured once and replayed. The graph outputs
5
+ option-letter logits for all positions [B, T, K]; slots are gathered outside.
6
+ """
7
+ import time, torch, torch.nn.functional as F
8
+ from .model import DecisionModel, collate
9
+ from .prompt import build, MAX_OPTIONS
10
+
11
+ T_BUCKETS = [64, 128, 192, 256, 320, 384, 512, 640, 768, 1024, 1280, 1536, 2048]
12
+ B_BUCKETS = [1, 2, 4, 8, 16, 32, 64]
13
+
14
+
15
+ def _bucket(x, buckets):
16
+ for b in buckets:
17
+ if x <= b:
18
+ return b
19
+ return None
20
+
21
+
22
+ class Engine:
23
+ def __init__(self, path, device="cuda", dtype=torch.bfloat16, use_graphs=True, max_ctx_tokens=1536):
24
+ self.m = DecisionModel(path, dtype=dtype, grad_ckpt=False).to(device).eval()
25
+ self.tok = self.m.tok; self.dev = device; self.use_graphs = use_graphs; self.max_ctx = max_ctx_tokens
26
+ self.core, self.W = self.m.lm.model, self.m.lm.lm_head.weight[self.m.letters].detach().clone()
27
+ self.graphs = {} # (B, T) -> (static_ids, static_out, graph)
28
+ self.pool = torch.cuda.graph_pool_handle() if use_graphs else None
29
+ self.stats = dict(graph_captures=0, forwards=0)
30
+
31
+ @torch.no_grad()
32
+ def _fwd(self, ids):
33
+ h = self.core(input_ids=ids).last_hidden_state
34
+ return F.linear(h, self.W).float() # [B, T, K]
35
+
36
+ def _capture(self, B, T):
37
+ s_ids = torch.full((B, T), self.tok.pad_token_id, dtype=torch.long, device=self.dev)
38
+ st = torch.cuda.Stream(); st.wait_stream(torch.cuda.current_stream())
39
+ with torch.cuda.stream(st):
40
+ for _ in range(2): self._fwd(s_ids) # warm-up: triton autotune etc.
41
+ torch.cuda.current_stream().wait_stream(st)
42
+ g = torch.cuda.CUDAGraph()
43
+ with torch.cuda.graph(g, pool=self.pool):
44
+ s_out = self._fwd(s_ids)
45
+ self.stats["graph_captures"] += 1
46
+ return s_ids, s_out, g
47
+
48
+ @torch.no_grad()
49
+ def logits_all(self, ids):
50
+ """ids: [B, T] long on device (already right-padded to a bucket). Returns [B, T, K] float."""
51
+ B, T = ids.shape; self.stats["forwards"] += 1
52
+ if not self.use_graphs:
53
+ return self._fwd(ids)
54
+ key = (B, T)
55
+ if key not in self.graphs:
56
+ self.graphs[key] = self._capture(B, T)
57
+ s_ids, s_out, g = self.graphs[key]
58
+ s_ids.copy_(ids); g.replay()
59
+ return s_out
60
+
61
+ @torch.no_grad()
62
+ def score_items(self, items, temperature=1.0):
63
+ """items: list of dicts from prompt.build. Returns list of [n_q, MAX_OPTIONS] prob tensors (cpu)."""
64
+ Tmax = max(len(it["ids"]) for it in items)
65
+ T = _bucket(Tmax, T_BUCKETS) or Tmax; B = _bucket(len(items), B_BUCKETS) or len(items)
66
+ ids = torch.full((B, T), self.tok.pad_token_id, dtype=torch.long)
67
+ for b, it in enumerate(items):
68
+ ids[b, :len(it["ids"])] = torch.tensor(it["ids"])
69
+ out = self.logits_all(ids.to(self.dev, non_blocking=True))
70
+ res = []
71
+ ar = torch.arange(MAX_OPTIONS, device=self.dev)
72
+ for b, it in enumerate(items):
73
+ sl = torch.tensor(it["slots"], device=self.dev)
74
+ lg = out[b, sl] # [n_q, K]
75
+ nop = torch.tensor(it["nopts"], device=self.dev)
76
+ lg = lg.masked_fill(ar[None, :] >= nop[:, None], float("-inf"))
77
+ res.append(torch.softmax(lg / temperature, -1).cpu())
78
+ return res
79
+
80
+ def warmup(self, shapes=((1, 128), (1, 256), (1, 384), (1, 512), (8, 256), (8, 512), (32, 256), (32, 512))):
81
+ t = time.time()
82
+ for B, T in shapes:
83
+ self.logits_all(torch.full((B, T), self.tok.pad_token_id, dtype=torch.long, device=self.dev))
84
+ torch.cuda.synchronize(); return time.time() - t
85
+
86
+
87
+ if __name__ == "__main__":
88
+ import sys, random, numpy as np
89
+ from . import data as D
90
+ from .infer import Decider
91
+ path = sys.argv[1] if len(sys.argv) > 1 else "runs/r3_v2/model"
92
+ _, evals = D.load_cache("data/tasks.pkl")
93
+ eng = Engine(path)
94
+ rng = random.Random(0)
95
+ exs = evals["support_tickets"][:64] + evals["clinc_oos"][:64] + evals["race"][:32]
96
+ items = [build(e, eng.tok, rng, max_ctx_tokens=1536) for e in exs]
97
+ # correctness vs eager masked forward (DecisionModel.slot_logits)
98
+ ref = []
99
+ with torch.no_grad():
100
+ for i in range(0, len(items), 16):
101
+ b = collate(items[i:i + 16], eng.tok.pad_token_id)
102
+ lg = eng.m.slot_logits(b["input_ids"].cuda(), b["attention_mask"].cuda(), b["slot_idx"].cuda(), b["slot_batch"].cuda(), b["nopts"].cuda())
103
+ ref.append(torch.softmax(lg, -1).cpu())
104
+ ref = torch.cat(ref)
105
+ got = torch.cat(eng.score_items(items))
106
+ print(f"max |p_graph - p_eager| = {(ref - got).abs().max():.4f} over {len(ref)} questions; argmax agreement {(ref.argmax(1) == got.argmax(1)).float().mean():.4f}")
107
+ print(f"warmup capture of 8 buckets: {eng.warmup():.1f}s; captures so far {eng.stats['graph_captures']}")
108
+ # latency: single real requests
109
+ for name, pool in [("support_tickets", exs[:64]), ("clinc_oos", exs[64:128]), ("race", exs[128:])]:
110
+ its = [build(e, eng.tok, rng) for e in pool]
111
+ ts = []
112
+ for it in its[:40]:
113
+ torch.cuda.synchronize(); t = time.time(); eng.score_items([it]); torch.cuda.synchronize(); ts.append(time.time() - t)
114
+ ts = np.array(ts[5:]) * 1000
115
+ print(f"single request {name:16s}: p50 {np.median(ts):5.1f} ms p90 {np.percentile(ts, 90):5.1f} ms (avg {np.mean([len(i['ids']) for i in its]):.0f} tok, {len(its[0]['slots'])} q)")
116
+ for bs in (8, 32):
117
+ ts = []
118
+ for i in range(0, min(len(its), bs * 6), bs):
119
+ chunk = its[i:i + bs]
120
+ if len(chunk) < bs: break
121
+ torch.cuda.synchronize(); t = time.time(); eng.score_items(chunk); torch.cuda.synchronize(); ts.append(time.time() - t)
122
+ ts = np.array(ts[1:]) * 1000
123
+ print(f" batch {bs:2d}: p50 {np.median(ts):6.1f} ms -> {bs/np.median(ts)*1000:6.0f} ctx/s, {bs*len(its[0]['slots'])/np.median(ts)*1000:6.0f} decisions/s")
124
+ print("stats", eng.stats, "graphs", len(eng.graphs), f"mem {torch.cuda.memory_reserved()/1e9:.1f} GB")
decider/infer.py ADDED
@@ -0,0 +1,145 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Usable inference API: typed decisions with probabilities, all from one forward pass.
2
+
3
+ from decider.infer import Decider
4
+ d = Decider("runs/r2_full/model")
5
+ out = d.decide("My card was charged twice for the same purchase.",
6
+ [{"question": "Which department should handle this?", "options": ["billing", "technical", "sales"]},
7
+ {"question": "How urgent is this?", "options": ["low", "medium", "high"]}])
8
+ # -> [{'choice': 'billing', 'confidence': 0.97, 'probs': {...}}, {...}]
9
+ """
10
+ import torch
11
+ from .model import DecisionModel, collate
12
+ from .prompt import build, MAX_OPTIONS
13
+ from dataclasses import dataclass
14
+
15
+
16
+ @dataclass
17
+ class Q:
18
+ text: str; options: list; gold: int = 0
19
+
20
+
21
+ @dataclass
22
+ class Example:
23
+ context: str; qs: list; task: str = "infer"
24
+
25
+
26
+ class Decider:
27
+ def __init__(self, path, device="cuda", dtype=torch.bfloat16, temperature=1.0, abstain_below=0.0):
28
+ self.m = DecisionModel(path, dtype=dtype, grad_ckpt=False).to(device).eval()
29
+ self.dev = device; self.T = temperature; self.abstain_below = abstain_below
30
+
31
+ @torch.no_grad()
32
+ def decide_batch(self, requests, max_ctx_tokens=1536):
33
+ """requests: list of (context:str, questions:list[dict(question, options)]). One forward pass for everything."""
34
+ exs, meta = [], []
35
+ for context, qs in requests:
36
+ for q in qs:
37
+ assert 2 <= len(q["options"]) <= MAX_OPTIONS, f"2..{MAX_OPTIONS} options required"
38
+ exs.append(Example(context, [Q(q["question"], list(q["options"]), 0) for q in qs], "infer"))
39
+ class _NoShuffle: # keep option order as given
40
+ def shuffle(self, x): pass
41
+ def sample(self, xs, k): return xs[:k]
42
+ items = [build(e, self.m.tok, _NoShuffle(), max_ctx_tokens=max_ctx_tokens) for e in exs]
43
+ b = collate(items, self.m.tok.pad_token_id)
44
+ logits = self.m.slot_logits(b["input_ids"].to(self.dev), b["attention_mask"].to(self.dev), b["slot_idx"].to(self.dev),
45
+ b["slot_batch"].to(self.dev), b["nopts"].to(self.dev))
46
+ probs = torch.softmax(logits / self.T, -1).cpu()
47
+ out, k = [], 0
48
+ for context, qs in requests:
49
+ res = []
50
+ for q in qs:
51
+ p = probs[k, :len(q["options"])].tolist(); k += 1
52
+ j = max(range(len(p)), key=p.__getitem__)
53
+ res.append(dict(choice=q["options"][j] if p[j] >= self.abstain_below else None, confidence=p[j],
54
+ probs={o: pi for o, pi in zip(q["options"], p)}))
55
+ out.append(res)
56
+ return out
57
+
58
+ def decide(self, context, questions, **kw):
59
+ return self.decide_batch([(context, questions)], **kw)[0]
60
+
61
+ # ---- typed schema interface: {question: {"type": "bool"} | {"type": "choice", "options": [...]}
62
+ # | {"type": "scale", "legend": {"0": "none", "1": "low", ...}}}
63
+ @staticmethod
64
+ def _schema_to_questions(schema):
65
+ qs = []
66
+ for qtext, spec in schema.items():
67
+ t = spec.get("type", "choice")
68
+ if t == "bool":
69
+ qs.append(dict(question=qtext, options=["no", "yes"]))
70
+ elif t == "choice":
71
+ qs.append(dict(question=qtext, options=list(spec["options"])))
72
+ elif t == "scale":
73
+ leg = spec["legend"]
74
+ keys = sorted(leg, key=lambda k: float(k)) if isinstance(leg, dict) else list(range(len(leg)))
75
+ labels = [f"{k}: {leg[k]}" if isinstance(leg, dict) else f"{i}: {leg[i]}" for i, k in enumerate(keys)]
76
+ qs.append(dict(question=qtext, options=labels, _keys=keys, _legend=leg))
77
+ else:
78
+ raise ValueError(f"unknown field type {t}")
79
+ return qs
80
+
81
+ def decide_json_batch(self, requests, **kw):
82
+ """requests: list of (context, schema). Returns one dict per context keyed by question."""
83
+ qss = [self._schema_to_questions(schema) for _, schema in requests]
84
+ raw = self.decide_batch([(ctx, qs) for (ctx, _), qs in zip(requests, qss)], **kw)
85
+ out = []
86
+ for (ctx, schema), qs, res in zip(requests, qss, raw):
87
+ o = {}
88
+ for (qtext, spec), q, r in zip(schema.items(), qs, res):
89
+ t = spec.get("type", "choice")
90
+ if t == "bool":
91
+ o[qtext] = {"noul": round(r["probs"]["yes"], 4), "type": "noul"}
92
+ elif t == "choice":
93
+ o[qtext] = {"choice": r["choice"], "confidence": round(r["confidence"], 4), "type": "choice",
94
+ "probabilities": {k: round(v, 4) for k, v in r["probs"].items()}}
95
+ else:
96
+ p = [r["probs"][lab] for lab in q["options"]]
97
+ keys = q["_keys"]; n = len(p)
98
+ score = sum(float(k) * pi for k, pi in zip(keys, p)) # expected level on the legend scale
99
+ j = max(range(n), key=p.__getitem__)
100
+ o[qtext] = {"score": round(score, 2), "confidence": round(p[j], 4), "type": "scale", "legend": q["_legend"],
101
+ "probabilities": {str(keys[i]): round(pi, 4) for i, pi in enumerate(p)}}
102
+ out.append(o)
103
+ return out
104
+
105
+ def decide_json(self, context, schema, **kw):
106
+ return self.decide_json_batch([(context, schema)], **kw)[0]
107
+
108
+
109
+ if __name__ == "__main__":
110
+ import sys, json, time
111
+ d = Decider(sys.argv[1] if len(sys.argv) > 1 else "runs/r1_200k/model")
112
+ demo = [
113
+ ("My card was charged twice for the same purchase and I want the extra charge refunded.",
114
+ [{"question": "Which department should handle this?", "options": ["billing", "technical support", "sales"]},
115
+ {"question": "What is the customer's sentiment?", "options": ["angry", "neutral", "happy"]},
116
+ {"question": "Does this need a refund action?", "options": ["no", "yes"]}]),
117
+ ("hey can u turn the lights off in the kitchen",
118
+ [{"question": "What is the intent?", "options": ["smart home control", "set alarm", "play music", "none of the above"]},
119
+ {"question": "Is this request toxic?", "options": ["no", "yes"]}]),
120
+ ("The quarterly report shows revenue fell 12% while costs rose sharply.",
121
+ [{"question": "What is the financial sentiment?", "options": ["bearish", "neutral", "bullish"]}]),
122
+ ]
123
+ t = time.time(); res = d.decide_batch(demo); dt = time.time() - t
124
+ for (ctx, qs), r in zip(demo, res):
125
+ print("\n>>", ctx)
126
+ for q, a in zip(qs, r):
127
+ print(f" {q['question']:45s} -> {a['choice']!s:22s} p={a['confidence']:.2f} " + " ".join(f"{o}:{p:.2f}" for o, p in a['probs'].items()))
128
+ print(f"\n{sum(len(q) for _, q in demo)} decisions in {dt*1000:.0f} ms (one forward pass)")
129
+ schema = {
130
+ "Revenue currently impacted?": {"type": "bool"},
131
+ "What business impact?": {"type": "choice", "options": ["none", "degraded", "outage"]},
132
+ "Integration issue present?": {"type": "bool"},
133
+ "Account health status?": {"type": "choice", "options": ["healthy", "watch", "at risk"]},
134
+ "Which incident scope?": {"type": "choice", "options": ["single_account", "multi_account", "platform_wide"]},
135
+ "Security concern present?": {"type": "bool"},
136
+ "Duplicate charge reported?": {"type": "bool"},
137
+ "Churn likelihood level?": {"type": "scale", "legend": {"0": "none", "1": "low", "2": "medium", "3": "high"}},
138
+ "Human attention needed?": {"type": "bool"},
139
+ "Immediate feature request?": {"type": "bool"},
140
+ }
141
+ ctx = ("Hi, since this morning our Stripe webhook integration stopped firing and our checkout is down for all customers. "
142
+ "We are losing orders every minute and our partner launch is on Thursday. Also I think we got billed twice last week. "
143
+ "If this is not fixed today we will have to look at other providers.")
144
+ t = time.time(); js = d.decide_json(ctx, schema); dt = time.time() - t
145
+ print(f"\n>> {ctx[:80]}...\n" + json.dumps(js, indent=1)[:3000]); print(f"{len(schema)} typed fields in {dt*1000:.0f} ms (one forward pass)")
decider/model.py ADDED
@@ -0,0 +1,47 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Backbone -> slot hidden states -> restricted logits over option letters."""
2
+ import torch, torch.nn as nn, torch.nn.functional as F
3
+ from transformers import AutoModelForCausalLM, AutoTokenizer
4
+ from .prompt import letter_ids, MAX_OPTIONS
5
+
6
+
7
+ class DecisionModel(nn.Module):
8
+ def __init__(self, name, dtype=torch.bfloat16, grad_ckpt=True):
9
+ super().__init__()
10
+ self.tok = AutoTokenizer.from_pretrained(name)
11
+ self.lm = AutoModelForCausalLM.from_pretrained(name, dtype=dtype)
12
+ if grad_ckpt:
13
+ self.lm.gradient_checkpointing_enable()
14
+ self.register_buffer("letters", torch.tensor(letter_ids(self.tok)), persistent=False)
15
+
16
+ def slot_logits(self, input_ids, attention_mask, slot_idx, slot_batch, nopts):
17
+ """input_ids [B,T]; slot_idx/slot_batch [N] flat slot positions; nopts [N].
18
+ Returns [N, MAX_OPTIONS] logits with invalid options masked to -inf."""
19
+ h = self.lm.model(input_ids=input_ids, attention_mask=attention_mask).last_hidden_state
20
+ hs = h[slot_batch, slot_idx] # [N,H]
21
+ W = self.lm.lm_head.weight[self.letters] # [K,H]
22
+ logits = F.linear(hs, W).float() # [N,K]
23
+ ar = torch.arange(MAX_OPTIONS, device=logits.device)[None, :]
24
+ logits = logits.masked_fill(ar >= nopts[:, None], float("-inf"))
25
+ return logits
26
+
27
+ def forward(self, batch):
28
+ return self.slot_logits(batch["input_ids"], batch["attention_mask"], batch["slot_idx"], batch["slot_batch"], batch["nopts"])
29
+
30
+
31
+ def collate(items, pad_id):
32
+ """items: list of dicts from prompt.build (+ 'task', 'ex_id'). Right-pad."""
33
+ T = max(len(it["ids"]) for it in items)
34
+ T = ((T + 63) // 64) * 64 # few distinct shapes -> fewer kernel (re)compiles
35
+ B = len(items)
36
+ input_ids = torch.full((B, T), pad_id, dtype=torch.long)
37
+ attn = torch.zeros((B, T), dtype=torch.long)
38
+ slot_idx, slot_batch, golds, nopts, tasks, qidx = [], [], [], [], [], []
39
+ for b, it in enumerate(items):
40
+ n = len(it["ids"])
41
+ input_ids[b, :n] = torch.tensor(it["ids"])
42
+ attn[b, :n] = 1
43
+ for k, s in enumerate(it["slots"]):
44
+ slot_idx.append(s); slot_batch.append(b); golds.append(it["golds"][k]); nopts.append(it["nopts"][k])
45
+ tasks.append(it.get("task", "")); qidx.append(k)
46
+ return dict(input_ids=input_ids, attention_mask=attn, slot_idx=torch.tensor(slot_idx), slot_batch=torch.tensor(slot_batch),
47
+ golds=torch.tensor(golds), nopts=torch.tensor(nopts), tasks=tasks, qidx=qidx)
decider/prompt.py ADDED
@@ -0,0 +1,52 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Prompt construction. One context, N typed questions, N answer slots.
2
+
3
+ All N decisions are read from a single forward pass: the logits at each
4
+ "Answer k: (" slot are restricted to the option-letter tokens. No answer
5
+ letters are ever inserted, so slot k sees the context and all questions but
6
+ no earlier answers (the decisions are conditionally independent given input).
7
+ """
8
+ import random
9
+
10
+ LETTERS = "ABCDEFGHIJ"
11
+ MAX_OPTIONS = len(LETTERS)
12
+
13
+
14
+ def build(example, tok, rng=None, max_options=MAX_OPTIONS, max_ctx_tokens=1536):
15
+ """Returns dict(ids=list[int], slots=list[int], golds=list[int], nopts=list[int], perms=list[list[int]])."""
16
+ rng = rng or random
17
+ ctx_ids = tok.encode("Context:\n" + example.context, add_special_tokens=False)[:max_ctx_tokens]
18
+ ids = list(ctx_ids)
19
+ slots, golds, nopts, perms = [], [], [], []
20
+ multi = len(example.qs) > 1
21
+ for k, q in enumerate(example.qs):
22
+ opts = list(range(len(q.options)))
23
+ if len(opts) > max_options:
24
+ others = [i for i in opts if i != q.gold]
25
+ keep = rng.sample(others, max_options - 1) + [q.gold]
26
+ opts = keep
27
+ rng.shuffle(opts)
28
+ lines = [f"\n\nQuestion{' ' + str(k + 1) if multi else ''}: {q.text}\nOptions:"]
29
+ for j, oi in enumerate(opts):
30
+ lines.append(f"\n({LETTERS[j]}) {q.options[oi]}")
31
+ lines.append(f"\nAnswer{' ' + str(k + 1) if multi else ''}: (")
32
+ piece = tok.encode("".join(lines), add_special_tokens=False)
33
+ ids.extend(piece)
34
+ slots.append(len(ids) - 1) # position of " (" token
35
+ golds.append(opts.index(q.gold) if q.gold in opts else -1)
36
+ nopts.append(len(opts))
37
+ perms.append(opts)
38
+ return dict(ids=ids, slots=slots, golds=golds, nopts=nopts, perms=perms)
39
+
40
+
41
+ def letter_ids(tok):
42
+ out = []
43
+ for L in LETTERS:
44
+ t = tok.encode(L, add_special_tokens=False)
45
+ assert len(t) == 1, (L, t)
46
+ out.append(t[0])
47
+ return out
48
+
49
+
50
+ def render(example, tok, **kw):
51
+ b = build(example, tok, **kw)
52
+ return tok.decode(b["ids"])
eval_results.json ADDED
@@ -0,0 +1,1283 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "results": {
3
+ "clinc_oos": {
4
+ "n": 1500,
5
+ "acc": 0.9766666666666667,
6
+ "nll": 0.07261181622743607,
7
+ "brier": 0.03852468356490135,
8
+ "ece": 0.010659157812595393,
9
+ "aurc": 0.0013287162793807251,
10
+ "acc_at_80": 0.9991666666666666,
11
+ "acc_at_50": 1.0,
12
+ "chance": 0.1,
13
+ "mean_conf": 0.9788497686386108,
14
+ "sec": 3.5,
15
+ "heldout": false
16
+ },
17
+ "massive_intent": {
18
+ "n": 1500,
19
+ "acc": 0.9493333333333334,
20
+ "nll": 0.15807873010635376,
21
+ "brier": 0.07667268067598343,
22
+ "ece": 0.012617475867271442,
23
+ "aurc": 0.005890076997496865,
24
+ "acc_at_80": 0.9933333333333333,
25
+ "acc_at_50": 0.9986666666666667,
26
+ "chance": 0.1,
27
+ "mean_conf": 0.9502438306808472,
28
+ "sec": 2.5,
29
+ "heldout": false
30
+ },
31
+ "massive_scenario": {
32
+ "n": 1500,
33
+ "acc": 0.8146666666666667,
34
+ "nll": 0.5371937155723572,
35
+ "brier": 0.2577172517776489,
36
+ "ece": 0.023450730780760452,
37
+ "aurc": 0.049948004214958666,
38
+ "acc_at_80": 0.9033333333333333,
39
+ "acc_at_50": 0.9706666666666667,
40
+ "chance": 0.1,
41
+ "mean_conf": 0.8208412528038025,
42
+ "sec": 2.4,
43
+ "heldout": true
44
+ },
45
+ "bitext_support": {
46
+ "n": 3000,
47
+ "acc": 0.9973333333333333,
48
+ "nll": 0.006749222986400127,
49
+ "brier": 0.0037496667355298996,
50
+ "ece": 0.0017460443178812545,
51
+ "aurc": 1.1261882856520142e-05,
52
+ "acc_at_80": 1.0,
53
+ "acc_at_50": 1.0,
54
+ "chance": 0.09999999999999998,
55
+ "mean_conf": 0.9978217482566833,
56
+ "sec": 3.4,
57
+ "heldout": false,
58
+ "per_q": [
59
+ 0.998,
60
+ 0.9966666666666667
61
+ ]
62
+ },
63
+ "support_tickets": {
64
+ "n": 4500,
65
+ "acc": 0.5791111111111111,
66
+ "nll": 0.9825455546379089,
67
+ "brier": 0.5169190764427185,
68
+ "ece": 0.016227112455500483,
69
+ "aurc": 0.21513268536739544,
70
+ "acc_at_80": 0.6405555555555555,
71
+ "acc_at_50": 0.7631111111111111,
72
+ "chance": 0.22777777777777777,
73
+ "mean_conf": 0.573243260383606,
74
+ "sec": 4.7,
75
+ "heldout": false,
76
+ "per_q": [
77
+ 0.8153333333333334,
78
+ 0.44,
79
+ 0.482
80
+ ]
81
+ },
82
+ "ag_news": {
83
+ "n": 1500,
84
+ "acc": 0.9173333333333333,
85
+ "nll": 0.24047400057315826,
86
+ "brier": 0.1257070153951645,
87
+ "ece": 0.010078805923461939,
88
+ "aurc": 0.018073040554543826,
89
+ "acc_at_80": 0.9716666666666667,
90
+ "acc_at_50": 0.988,
91
+ "chance": 0.25,
92
+ "mean_conf": 0.914076030254364,
93
+ "sec": 2.5,
94
+ "heldout": false
95
+ },
96
+ "dbpedia": {
97
+ "n": 1500,
98
+ "acc": 0.9913333333333333,
99
+ "nll": 0.0235174261033535,
100
+ "brier": 0.012497784569859505,
101
+ "ece": 0.0035786836942037336,
102
+ "aurc": 0.00010865530954364588,
103
+ "acc_at_80": 1.0,
104
+ "acc_at_50": 1.0,
105
+ "chance": 0.1,
106
+ "mean_conf": 0.9918569326400757,
107
+ "sec": 3.2,
108
+ "heldout": false
109
+ },
110
+ "yahoo_topics": {
111
+ "n": 1500,
112
+ "acc": 0.762,
113
+ "nll": 0.7322853207588196,
114
+ "brier": 0.33554312586784363,
115
+ "ece": 0.024258904794851946,
116
+ "aurc": 0.08055392547857332,
117
+ "acc_at_80": 0.8533333333333334,
118
+ "acc_at_50": 0.9426666666666667,
119
+ "chance": 0.1,
120
+ "mean_conf": 0.7577335834503174,
121
+ "sec": 2.9,
122
+ "heldout": false
123
+ },
124
+ "newsgroups": {
125
+ "n": 1462,
126
+ "acc": 0.8146374829001368,
127
+ "nll": 0.5313005447387695,
128
+ "brier": 0.25277355313301086,
129
+ "ece": 0.02700514194495698,
130
+ "aurc": 0.03996642898945821,
131
+ "acc_at_80": 0.917094017094017,
132
+ "acc_at_50": 0.9835841313269493,
133
+ "chance": 0.09999999999999999,
134
+ "mean_conf": 0.8385738730430603,
135
+ "sec": 5.3,
136
+ "heldout": false
137
+ },
138
+ "bbc_news": {
139
+ "n": 1000,
140
+ "acc": 0.941,
141
+ "nll": 0.15981632471084595,
142
+ "brier": 0.09309598803520203,
143
+ "ece": 0.01646793046593665,
144
+ "aurc": 0.007680621955810168,
145
+ "acc_at_80": 0.98375,
146
+ "acc_at_50": 1.0,
147
+ "chance": 0.20000000000000004,
148
+ "mean_conf": 0.936293363571167,
149
+ "sec": 5.3,
150
+ "heldout": true
151
+ },
152
+ "trec": {
153
+ "n": 500,
154
+ "acc": 0.766,
155
+ "nll": 0.5928767323493958,
156
+ "brier": 0.31943121552467346,
157
+ "ece": 0.03574929744005202,
158
+ "aurc": 0.09966853443922356,
159
+ "acc_at_80": 0.8275,
160
+ "acc_at_50": 0.916,
161
+ "chance": 0.16666666666666663,
162
+ "mean_conf": 0.769515872001648,
163
+ "sec": 0.9,
164
+ "heldout": true
165
+ },
166
+ "student_questions": {
167
+ "n": 1500,
168
+ "acc": 0.9273333333333333,
169
+ "nll": 0.22138677537441254,
170
+ "brier": 0.11496378481388092,
171
+ "ece": 0.029807024538516993,
172
+ "aurc": 0.015550148490194371,
173
+ "acc_at_80": 0.9716666666666667,
174
+ "acc_at_50": 0.992,
175
+ "chance": 0.25,
176
+ "mean_conf": 0.9056621193885803,
177
+ "sec": 3.0,
178
+ "heldout": true
179
+ },
180
+ "dolly_category": {
181
+ "n": 1500,
182
+ "acc": 0.318,
183
+ "nll": 2.0031864643096924,
184
+ "brier": 0.8271570801734924,
185
+ "ece": 0.21146952118476237,
186
+ "aurc": 0.48668464321903865,
187
+ "acc_at_80": 0.3575,
188
+ "acc_at_50": 0.44,
189
+ "chance": 0.125,
190
+ "mean_conf": 0.5294612646102905,
191
+ "sec": 3.5,
192
+ "heldout": true
193
+ },
194
+ "imdb": {
195
+ "n": 1500,
196
+ "acc": 0.9633333333333334,
197
+ "nll": 0.0927785336971283,
198
+ "brier": 0.053890664130449295,
199
+ "ece": 0.007267733812332158,
200
+ "aurc": 0.0027718147959179054,
201
+ "acc_at_80": 0.9983333333333333,
202
+ "acc_at_50": 1.0,
203
+ "chance": 0.5,
204
+ "mean_conf": 0.9614710211753845,
205
+ "sec": 5.6,
206
+ "heldout": false
207
+ },
208
+ "sst2": {
209
+ "n": 872,
210
+ "acc": 0.9415137614678899,
211
+ "nll": 0.15022407472133636,
212
+ "brier": 0.08405344188213348,
213
+ "ece": 0.014767689784185644,
214
+ "aurc": 0.009848323778768498,
215
+ "acc_at_80": 0.9899713467048711,
216
+ "acc_at_50": 0.9931192660550459,
217
+ "chance": 0.5,
218
+ "mean_conf": 0.9442880749702454,
219
+ "sec": 1.5,
220
+ "heldout": false
221
+ },
222
+ "sst5": {
223
+ "n": 1500,
224
+ "acc": 0.6133333333333333,
225
+ "nll": 0.9277896881103516,
226
+ "brier": 0.5273030996322632,
227
+ "ece": 0.03280426164468128,
228
+ "aurc": 0.30489681515108674,
229
+ "acc_at_80": 0.6466666666666666,
230
+ "acc_at_50": 0.684,
231
+ "chance": 0.2,
232
+ "mean_conf": 0.58907151222229,
233
+ "sec": 2.4,
234
+ "heldout": false
235
+ },
236
+ "yelp": {
237
+ "n": 1500,
238
+ "acc": 0.69,
239
+ "nll": 0.7099072337150574,
240
+ "brier": 0.41327351331710815,
241
+ "ece": 0.023333197732766472,
242
+ "aurc": 0.16233091407117872,
243
+ "acc_at_80": 0.7483333333333333,
244
+ "acc_at_50": 0.8146666666666667,
245
+ "chance": 0.2,
246
+ "mean_conf": 0.7003123164176941,
247
+ "sec": 4.5,
248
+ "heldout": false
249
+ },
250
+ "amazon_stars": {
251
+ "n": 1500,
252
+ "acc": 0.6066666666666667,
253
+ "nll": 0.9025667309761047,
254
+ "brier": 0.5039957761764526,
255
+ "ece": 0.030678236802419018,
256
+ "aurc": 0.2318602277252905,
257
+ "acc_at_80": 0.6533333333333333,
258
+ "acc_at_50": 0.7546666666666667,
259
+ "chance": 0.2,
260
+ "mean_conf": 0.6180839538574219,
261
+ "sec": 2.7,
262
+ "heldout": false
263
+ },
264
+ "emotion": {
265
+ "n": 1500,
266
+ "acc": 0.854,
267
+ "nll": 0.3842936158180237,
268
+ "brier": 0.20495271682739258,
269
+ "ece": 0.014472351928551988,
270
+ "aurc": 0.03513619107373314,
271
+ "acc_at_80": 0.9308333333333333,
272
+ "acc_at_50": 0.9813333333333333,
273
+ "chance": 0.16666666666666666,
274
+ "mean_conf": 0.8586899638175964,
275
+ "sec": 2.4,
276
+ "heldout": false
277
+ },
278
+ "go_emotions": {
279
+ "n": 1500,
280
+ "acc": 0.7766666666666666,
281
+ "nll": 0.6451207399368286,
282
+ "brier": 0.3155723512172699,
283
+ "ece": 0.01979676757256191,
284
+ "aurc": 0.07662852818492483,
285
+ "acc_at_80": 0.8558333333333333,
286
+ "acc_at_50": 0.9453333333333334,
287
+ "chance": 0.1,
288
+ "mean_conf": 0.787257730960846,
289
+ "sec": 2.4,
290
+ "heldout": false
291
+ },
292
+ "tweet_sentiment": {
293
+ "n": 1500,
294
+ "acc": 0.7353333333333333,
295
+ "nll": 0.5786967873573303,
296
+ "brier": 0.35522446036338806,
297
+ "ece": 0.02083511827389399,
298
+ "aurc": 0.12570788661149288,
299
+ "acc_at_80": 0.7933333333333333,
300
+ "acc_at_50": 0.868,
301
+ "chance": 0.3333333333333333,
302
+ "mean_conf": 0.7195278406143188,
303
+ "sec": 2.4,
304
+ "heldout": false
305
+ },
306
+ "tweet_emotion": {
307
+ "n": 1421,
308
+ "acc": 0.8395496129486277,
309
+ "nll": 0.4268813133239746,
310
+ "brier": 0.22563911974430084,
311
+ "ece": 0.03929315151923309,
312
+ "aurc": 0.03901553887763414,
313
+ "acc_at_80": 0.9217238346525946,
314
+ "acc_at_50": 0.9830985915492958,
315
+ "chance": 0.25,
316
+ "mean_conf": 0.8121002912521362,
317
+ "sec": 2.3,
318
+ "heldout": false
319
+ },
320
+ "tweet_irony": {
321
+ "n": 784,
322
+ "acc": 0.7793367346938775,
323
+ "nll": 0.4821619391441345,
324
+ "brier": 0.3177129328250885,
325
+ "ece": 0.05867012658593607,
326
+ "aurc": 0.10753610876267371,
327
+ "acc_at_80": 0.8373205741626795,
328
+ "acc_at_50": 0.8928571428571429,
329
+ "chance": 0.5,
330
+ "mean_conf": 0.7206665277481079,
331
+ "sec": 1.3,
332
+ "heldout": true
333
+ },
334
+ "fin_sentiment": {
335
+ "n": 1500,
336
+ "acc": 0.7693333333333333,
337
+ "nll": 0.5099801421165466,
338
+ "brier": 0.313610315322876,
339
+ "ece": 0.03050452242294948,
340
+ "aurc": 0.0985651541028985,
341
+ "acc_at_80": 0.8283333333333334,
342
+ "acc_at_50": 0.9186666666666666,
343
+ "chance": 0.3333333333333333,
344
+ "mean_conf": 0.7750865817070007,
345
+ "sec": 2.4,
346
+ "heldout": true
347
+ },
348
+ "cr_reviews": {
349
+ "n": 753,
350
+ "acc": 0.9017264276228419,
351
+ "nll": 0.23763678967952728,
352
+ "brier": 0.14010722935199738,
353
+ "ece": 0.02007227114947191,
354
+ "aurc": 0.02070501416560806,
355
+ "acc_at_80": 0.9634551495016611,
356
+ "acc_at_50": 0.9946808510638298,
357
+ "chance": 0.5,
358
+ "mean_conf": 0.9198644161224365,
359
+ "sec": 1.3,
360
+ "heldout": true
361
+ },
362
+ "counterfactual": {
363
+ "n": 1500,
364
+ "acc": 0.952,
365
+ "nll": 0.12973414361476898,
366
+ "brier": 0.07501313090324402,
367
+ "ece": 0.009895903388659118,
368
+ "aurc": 0.0057674469796164165,
369
+ "acc_at_80": 0.9891666666666666,
370
+ "acc_at_50": 1.0,
371
+ "chance": 0.5,
372
+ "mean_conf": 0.9529626369476318,
373
+ "sec": 2.5,
374
+ "heldout": false
375
+ },
376
+ "subjectivity": {
377
+ "n": 1500,
378
+ "acc": 0.9653333333333334,
379
+ "nll": 0.1150687038898468,
380
+ "brier": 0.06236080080270767,
381
+ "ece": 0.014521381457646683,
382
+ "aurc": 0.005428128396492908,
383
+ "acc_at_80": 0.9916666666666667,
384
+ "acc_at_50": 0.9973333333333333,
385
+ "chance": 0.5,
386
+ "mean_conf": 0.9508118629455566,
387
+ "sec": 2.4,
388
+ "heldout": false
389
+ },
390
+ "tweet_offensive": {
391
+ "n": 860,
392
+ "acc": 0.8569767441860465,
393
+ "nll": 0.3296823799610138,
394
+ "brier": 0.2049010545015335,
395
+ "ece": 0.048828685491584076,
396
+ "aurc": 0.04312619584445423,
397
+ "acc_at_80": 0.9229651162790697,
398
+ "acc_at_50": 0.9674418604651163,
399
+ "chance": 0.5,
400
+ "mean_conf": 0.8316621780395508,
401
+ "sec": 1.4,
402
+ "heldout": false
403
+ },
404
+ "tweet_hate": {
405
+ "n": 1500,
406
+ "acc": 0.5126666666666667,
407
+ "nll": 0.8970687389373779,
408
+ "brier": 0.634926974773407,
409
+ "ece": 0.25461129041512803,
410
+ "aurc": 0.41152814993800513,
411
+ "acc_at_80": 0.535,
412
+ "acc_at_50": 0.5426666666666666,
413
+ "chance": 0.5,
414
+ "mean_conf": 0.7672780156135559,
415
+ "sec": 2.4,
416
+ "heldout": false
417
+ },
418
+ "hate_offensive": {
419
+ "n": 1500,
420
+ "acc": 0.9213333333333333,
421
+ "nll": 0.2372422218322754,
422
+ "brier": 0.12591436505317688,
423
+ "ece": 0.012463796754678066,
424
+ "aurc": 0.01799362902327372,
425
+ "acc_at_80": 0.9716666666666667,
426
+ "acc_at_50": 0.988,
427
+ "chance": 0.3333333333333333,
428
+ "mean_conf": 0.9196624159812927,
429
+ "sec": 2.4,
430
+ "heldout": false
431
+ },
432
+ "civil_comments": {
433
+ "n": 7500,
434
+ "acc": 0.9353333333333333,
435
+ "nll": 0.1511416733264923,
436
+ "brier": 0.09208963811397552,
437
+ "ece": 0.004735005791982015,
438
+ "aurc": 0.007750731444351848,
439
+ "acc_at_80": 0.9876666666666667,
440
+ "acc_at_50": 0.9997333333333334,
441
+ "chance": 0.5,
442
+ "mean_conf": 0.9398914575576782,
443
+ "sec": 4.2,
444
+ "heldout": false,
445
+ "per_q": [
446
+ 0.872,
447
+ 0.978,
448
+ 0.9906666666666667,
449
+ 0.8666666666666667,
450
+ 0.9693333333333334
451
+ ]
452
+ },
453
+ "toxic_chat": {
454
+ "n": 3000,
455
+ "acc": 0.979,
456
+ "nll": 0.055397775024175644,
457
+ "brier": 0.03014867566525936,
458
+ "ece": 0.004773223658402735,
459
+ "aurc": 0.0011957276641795629,
460
+ "acc_at_80": 0.99875,
461
+ "acc_at_50": 1.0,
462
+ "chance": 0.5,
463
+ "mean_conf": 0.9801945090293884,
464
+ "sec": 2.8,
465
+ "heldout": false,
466
+ "per_q": [
467
+ 0.9646666666666667,
468
+ 0.9933333333333333
469
+ ]
470
+ },
471
+ "sms_spam": {
472
+ "n": 1000,
473
+ "acc": 0.986,
474
+ "nll": 0.05686425045132637,
475
+ "brier": 0.025409400463104248,
476
+ "ece": 0.009181047439575173,
477
+ "aurc": 0.002507962440114501,
478
+ "acc_at_80": 0.99625,
479
+ "acc_at_50": 0.998,
480
+ "chance": 0.5,
481
+ "mean_conf": 0.989531397819519,
482
+ "sec": 1.7,
483
+ "heldout": false
484
+ },
485
+ "enron_spam": {
486
+ "n": 1500,
487
+ "acc": 0.9873333333333333,
488
+ "nll": 0.027067676186561584,
489
+ "brier": 0.016892286017537117,
490
+ "ece": 0.004601201574007648,
491
+ "aurc": 0.00021669720532009002,
492
+ "acc_at_80": 1.0,
493
+ "acc_at_50": 1.0,
494
+ "chance": 0.5,
495
+ "mean_conf": 0.9897943735122681,
496
+ "sec": 5.8,
497
+ "heldout": false
498
+ },
499
+ "insincere_questions": {
500
+ "n": 1500,
501
+ "acc": 0.9553333333333334,
502
+ "nll": 0.10724982619285583,
503
+ "brier": 0.06384952366352081,
504
+ "ece": 0.01187876927852626,
505
+ "aurc": 0.0043015761367758325,
506
+ "acc_at_80": 0.9975,
507
+ "acc_at_50": 0.9986666666666667,
508
+ "chance": 0.5,
509
+ "mean_conf": 0.9536961317062378,
510
+ "sec": 2.4,
511
+ "heldout": false
512
+ },
513
+ "ade": {
514
+ "n": 1500,
515
+ "acc": 0.808,
516
+ "nll": 0.44578930735588074,
517
+ "brier": 0.2858363091945648,
518
+ "ece": 0.03843762199083964,
519
+ "aurc": 0.09344845724457498,
520
+ "acc_at_80": 0.8516666666666667,
521
+ "acc_at_50": 0.9133333333333333,
522
+ "chance": 0.5,
523
+ "mean_conf": 0.8236058950424194,
524
+ "sec": 2.4,
525
+ "heldout": true
526
+ },
527
+ "snli": {
528
+ "n": 1476,
529
+ "acc": 0.9295392953929539,
530
+ "nll": 0.22098393738269806,
531
+ "brier": 0.11396971344947815,
532
+ "ece": 0.02313795467702355,
533
+ "aurc": 0.01651608891549282,
534
+ "acc_at_80": 0.9729043183742591,
535
+ "acc_at_50": 0.986449864498645,
536
+ "chance": 0.3333333333333333,
537
+ "mean_conf": 0.9117720127105713,
538
+ "sec": 2.4,
539
+ "heldout": false
540
+ },
541
+ "mnli": {
542
+ "n": 1500,
543
+ "acc": 0.8846666666666667,
544
+ "nll": 0.29387927055358887,
545
+ "brier": 0.16493351757526398,
546
+ "ece": 0.030576728582382185,
547
+ "aurc": 0.023789420553220787,
548
+ "acc_at_80": 0.9616666666666667,
549
+ "acc_at_50": 0.9906666666666667,
550
+ "chance": 0.3333333333333333,
551
+ "mean_conf": 0.8777571320533752,
552
+ "sec": 2.5,
553
+ "heldout": false
554
+ },
555
+ "rte": {
556
+ "n": 277,
557
+ "acc": 0.8592057761732852,
558
+ "nll": 0.3232092559337616,
559
+ "brier": 0.198373481631279,
560
+ "ece": 0.03932678484314188,
561
+ "aurc": 0.0427986576794893,
562
+ "acc_at_80": 0.918918918918919,
563
+ "acc_at_50": 0.9710144927536232,
564
+ "chance": 0.5,
565
+ "mean_conf": 0.8895930647850037,
566
+ "sec": 0.6,
567
+ "heldout": false
568
+ },
569
+ "qnli": {
570
+ "n": 1500,
571
+ "acc": 0.9393333333333334,
572
+ "nll": 0.15822429955005646,
573
+ "brier": 0.08866540342569351,
574
+ "ece": 0.01913584089279174,
575
+ "aurc": 0.009447146399598809,
576
+ "acc_at_80": 0.9841666666666666,
577
+ "acc_at_50": 0.9973333333333333,
578
+ "chance": 0.5,
579
+ "mean_conf": 0.9270872473716736,
580
+ "sec": 2.4,
581
+ "heldout": false
582
+ },
583
+ "qqp": {
584
+ "n": 1500,
585
+ "acc": 0.8533333333333334,
586
+ "nll": 0.3225935399532318,
587
+ "brier": 0.2073184996843338,
588
+ "ece": 0.010325272003809595,
589
+ "aurc": 0.04126678230090233,
590
+ "acc_at_80": 0.9141666666666667,
591
+ "acc_at_50": 0.9773333333333334,
592
+ "chance": 0.5,
593
+ "mean_conf": 0.8565818667411804,
594
+ "sec": 2.4,
595
+ "heldout": false
596
+ },
597
+ "mrpc": {
598
+ "n": 408,
599
+ "acc": 0.8455882352941176,
600
+ "nll": 0.33268094062805176,
601
+ "brier": 0.21494312584400177,
602
+ "ece": 0.052983781724583866,
603
+ "aurc": 0.04073823335434742,
604
+ "acc_at_80": 0.9079754601226994,
605
+ "acc_at_50": 0.9803921568627451,
606
+ "chance": 0.5,
607
+ "mean_conf": 0.8508865237236023,
608
+ "sec": 0.7,
609
+ "heldout": false
610
+ },
611
+ "paws": {
612
+ "n": 1500,
613
+ "acc": 0.6906666666666667,
614
+ "nll": 0.8107166290283203,
615
+ "brier": 0.4746285080909729,
616
+ "ece": 0.22073563802242277,
617
+ "aurc": 0.14589154286965242,
618
+ "acc_at_80": 0.765,
619
+ "acc_at_50": 0.8693333333333333,
620
+ "chance": 0.5,
621
+ "mean_conf": 0.9114023447036743,
622
+ "sec": 2.4,
623
+ "heldout": true
624
+ },
625
+ "cola": {
626
+ "n": 1043,
627
+ "acc": 0.8082454458293384,
628
+ "nll": 0.40937936305999756,
629
+ "brier": 0.2624674439430237,
630
+ "ece": 0.03032427043265601,
631
+ "aurc": 0.07498094574041442,
632
+ "acc_at_80": 0.8669064748201439,
633
+ "acc_at_50": 0.9367816091954023,
634
+ "chance": 0.5,
635
+ "mean_conf": 0.8143807053565979,
636
+ "sec": 1.7,
637
+ "heldout": false
638
+ },
639
+ "boolq": {
640
+ "n": 1500,
641
+ "acc": 0.884,
642
+ "nll": 0.28092503547668457,
643
+ "brier": 0.16896213591098785,
644
+ "ece": 0.013350882569948819,
645
+ "aurc": 0.03244594523564409,
646
+ "acc_at_80": 0.9425,
647
+ "acc_at_50": 0.9733333333333334,
648
+ "chance": 0.5,
649
+ "mean_conf": 0.8836175799369812,
650
+ "sec": 3.7,
651
+ "heldout": false
652
+ },
653
+ "strategyqa": {
654
+ "n": 687,
655
+ "acc": 0.6128093158660844,
656
+ "nll": 0.6776027679443359,
657
+ "brier": 0.4776630997657776,
658
+ "ece": 0.0688689137650369,
659
+ "aurc": 0.31974077928206984,
660
+ "acc_at_80": 0.6272727272727273,
661
+ "acc_at_50": 0.6656976744186046,
662
+ "chance": 0.5,
663
+ "mean_conf": 0.6726659536361694,
664
+ "sec": 1.2,
665
+ "heldout": true
666
+ },
667
+ "pubmedqa": {
668
+ "n": 500,
669
+ "acc": 0.79,
670
+ "nll": 0.5823958516120911,
671
+ "brier": 0.3185318112373352,
672
+ "ece": 0.06283757454156877,
673
+ "aurc": 0.09994473800171019,
674
+ "acc_at_80": 0.8525,
675
+ "acc_at_50": 0.924,
676
+ "chance": 0.33333333333333326,
677
+ "mean_conf": 0.7943322658538818,
678
+ "sec": 2.4,
679
+ "heldout": true
680
+ },
681
+ "arc": {
682
+ "n": 1172,
683
+ "acc": 0.8575085324232082,
684
+ "nll": 0.37918952107429504,
685
+ "brier": 0.19906748831272125,
686
+ "ece": 0.029848500348805575,
687
+ "aurc": 0.028284304075720537,
688
+ "acc_at_80": 0.9477611940298507,
689
+ "acc_at_50": 0.9880546075085325,
690
+ "chance": 0.25015642775881686,
691
+ "mean_conf": 0.8538938164710999,
692
+ "sec": 1.9,
693
+ "heldout": false
694
+ },
695
+ "commonsense_qa": {
696
+ "n": 1221,
697
+ "acc": 0.7592137592137592,
698
+ "nll": 0.5952430367469788,
699
+ "brier": 0.31855452060699463,
700
+ "ece": 0.030964655577404226,
701
+ "aurc": 0.07921845119739539,
702
+ "acc_at_80": 0.8382804503582395,
703
+ "acc_at_50": 0.9475409836065574,
704
+ "chance": 0.19999999999999998,
705
+ "mean_conf": 0.7593477368354797,
706
+ "sec": 2.0,
707
+ "heldout": false
708
+ },
709
+ "qasc": {
710
+ "n": 926,
711
+ "acc": 0.816414686825054,
712
+ "nll": 0.5963356494903564,
713
+ "brier": 0.27422401309013367,
714
+ "ece": 0.07903238865456365,
715
+ "aurc": 0.04646154243439387,
716
+ "acc_at_80": 0.9149797570850202,
717
+ "acc_at_50": 0.978401727861771,
718
+ "chance": 0.125,
719
+ "mean_conf": 0.7387062311172485,
720
+ "sec": 1.5,
721
+ "heldout": false
722
+ },
723
+ "openbookqa": {
724
+ "n": 500,
725
+ "acc": 0.824,
726
+ "nll": 0.4846096932888031,
727
+ "brier": 0.250370055437088,
728
+ "ece": 0.04813598054647445,
729
+ "aurc": 0.05417227716258954,
730
+ "acc_at_80": 0.905,
731
+ "acc_at_50": 0.964,
732
+ "chance": 0.25,
733
+ "mean_conf": 0.7986846566200256,
734
+ "sec": 0.8,
735
+ "heldout": false
736
+ },
737
+ "sciq": {
738
+ "n": 1000,
739
+ "acc": 0.982,
740
+ "nll": 0.05202531814575195,
741
+ "brier": 0.025178927928209305,
742
+ "ece": 0.017304384350776688,
743
+ "aurc": 0.000406837269511718,
744
+ "acc_at_80": 1.0,
745
+ "acc_at_50": 1.0,
746
+ "chance": 0.25,
747
+ "mean_conf": 0.9719231128692627,
748
+ "sec": 2.3,
749
+ "heldout": true
750
+ },
751
+ "hellaswag": {
752
+ "n": 1500,
753
+ "acc": 0.8733333333333333,
754
+ "nll": 0.3213328719139099,
755
+ "brier": 0.17381860315799713,
756
+ "ece": 0.027953611870606753,
757
+ "aurc": 0.021264232574898293,
758
+ "acc_at_80": 0.9566666666666667,
759
+ "acc_at_50": 0.9986666666666667,
760
+ "chance": 0.25,
761
+ "mean_conf": 0.8975942134857178,
762
+ "sec": 4.3,
763
+ "heldout": false
764
+ },
765
+ "piqa": {
766
+ "n": 1500,
767
+ "acc": 0.8366666666666667,
768
+ "nll": 0.35584479570388794,
769
+ "brier": 0.22590981423854828,
770
+ "ece": 0.021878054340680464,
771
+ "aurc": 0.052599671905179834,
772
+ "acc_at_80": 0.9025,
773
+ "acc_at_50": 0.956,
774
+ "chance": 0.5,
775
+ "mean_conf": 0.8382094502449036,
776
+ "sec": 2.6,
777
+ "heldout": false
778
+ },
779
+ "social_iqa": {
780
+ "n": 1500,
781
+ "acc": 0.712,
782
+ "nll": 0.6847297549247742,
783
+ "brier": 0.3952922224998474,
784
+ "ece": 0.060257167994976046,
785
+ "aurc": 0.14267195194658805,
786
+ "acc_at_80": 0.7783333333333333,
787
+ "acc_at_50": 0.8493333333333334,
788
+ "chance": 0.3333333333333333,
789
+ "mean_conf": 0.7657021284103394,
790
+ "sec": 2.4,
791
+ "heldout": true
792
+ },
793
+ "winogrande": {
794
+ "n": 1267,
795
+ "acc": 0.6748224151539068,
796
+ "nll": 0.6599352359771729,
797
+ "brier": 0.43405067920684814,
798
+ "ece": 0.12888923459101967,
799
+ "aurc": 0.19936165776962148,
800
+ "acc_at_80": 0.7199211045364892,
801
+ "acc_at_50": 0.8012618296529969,
802
+ "chance": 0.5,
803
+ "mean_conf": 0.8037116527557373,
804
+ "sec": 2.1,
805
+ "heldout": false
806
+ },
807
+ "race": {
808
+ "n": 1500,
809
+ "acc": 0.8406666666666667,
810
+ "nll": 0.45064714550971985,
811
+ "brier": 0.22838135063648224,
812
+ "ece": 0.024594745457172388,
813
+ "aurc": 0.04331003810864532,
814
+ "acc_at_80": 0.92,
815
+ "acc_at_50": 0.9773333333333334,
816
+ "chance": 0.25,
817
+ "mean_conf": 0.8616733551025391,
818
+ "sec": 7.7,
819
+ "heldout": false
820
+ },
821
+ "mmlu": {
822
+ "n": 1500,
823
+ "acc": 0.632,
824
+ "nll": 0.9076293110847473,
825
+ "brier": 0.46681055426597595,
826
+ "ece": 0.047448410332202914,
827
+ "aurc": 0.1664288415331668,
828
+ "acc_at_80": 0.7158333333333333,
829
+ "acc_at_50": 0.8466666666666667,
830
+ "chance": 0.25,
831
+ "mean_conf": 0.6779200434684753,
832
+ "sec": 3.1,
833
+ "heldout": false
834
+ },
835
+ "medqa": {
836
+ "n": 1273,
837
+ "acc": 0.5640219952867243,
838
+ "nll": 1.0263904333114624,
839
+ "brier": 0.5558311939239502,
840
+ "ece": 0.024355607152640676,
841
+ "aurc": 0.26470923464312796,
842
+ "acc_at_80": 0.6198428290766208,
843
+ "acc_at_50": 0.7044025157232704,
844
+ "chance": 0.25,
845
+ "mean_conf": 0.5750094056129456,
846
+ "sec": 4.0,
847
+ "heldout": false
848
+ },
849
+ "truthfulqa": {
850
+ "n": 817,
851
+ "acc": 0.4969400244798042,
852
+ "nll": 1.4452694654464722,
853
+ "brier": 0.6520512104034424,
854
+ "ece": 0.09137629450388902,
855
+ "aurc": 0.29796415629912126,
856
+ "acc_at_80": 0.5565749235474006,
857
+ "acc_at_50": 0.6764705882352942,
858
+ "chance": 0.22622350449767833,
859
+ "mean_conf": 0.5813891291618347,
860
+ "sec": 1.4,
861
+ "heldout": true
862
+ },
863
+ "fin_phrasebank": {
864
+ "n": 970,
865
+ "acc": 0.6597938144329897,
866
+ "nll": 0.6683916449546814,
867
+ "brier": 0.42705729603767395,
868
+ "ece": 0.08229427663321345,
869
+ "aurc": 0.1944303228748502,
870
+ "acc_at_80": 0.7190721649484536,
871
+ "acc_at_50": 0.8123711340206186,
872
+ "chance": 0.33333333333333326,
873
+ "mean_conf": 0.7379865050315857,
874
+ "sec": 4.0,
875
+ "heldout": true
876
+ },
877
+ "banking77": {
878
+ "n": 1500,
879
+ "acc": 0.96,
880
+ "nll": 0.12018896639347076,
881
+ "brier": 0.058943431824445724,
882
+ "ece": 0.010675735374291784,
883
+ "aurc": 0.0028384511661530853,
884
+ "acc_at_80": 0.9966666666666667,
885
+ "acc_at_50": 1.0,
886
+ "chance": 0.1,
887
+ "mean_conf": 0.9642454981803894,
888
+ "sec": 2.4,
889
+ "heldout": false
890
+ },
891
+ "bias_in_bios": {
892
+ "n": 3000,
893
+ "acc": 0.953,
894
+ "nll": 0.13723737001419067,
895
+ "brier": 0.06831751763820648,
896
+ "ece": 0.007673392459750155,
897
+ "aurc": 0.003360504536259007,
898
+ "acc_at_80": 0.9979166666666667,
899
+ "acc_at_50": 1.0,
900
+ "chance": 0.30000000000000004,
901
+ "mean_conf": 0.9521626234054565,
902
+ "sec": 3.9,
903
+ "heldout": false,
904
+ "per_q": [
905
+ 0.9086666666666666,
906
+ 0.9973333333333333
907
+ ]
908
+ },
909
+ "helpsteer2": {
910
+ "n": 5190,
911
+ "acc": 0.6021194605009634,
912
+ "nll": 0.9371625185012817,
913
+ "brier": 0.5185062289237976,
914
+ "ece": 0.018472507194057828,
915
+ "aurc": 0.26482025743101484,
916
+ "acc_at_80": 0.6572736030828517,
917
+ "acc_at_50": 0.735645472061657,
918
+ "chance": 0.2,
919
+ "mean_conf": 0.6043042540550232,
920
+ "sec": 9.6,
921
+ "heldout": false,
922
+ "per_q": [
923
+ 0.45857418111753373,
924
+ 0.5173410404624278,
925
+ 0.7244701348747592,
926
+ 0.6242774566473989,
927
+ 0.6859344894026975
928
+ ]
929
+ },
930
+ "helpsteer3_pref": {
931
+ "n": 1176,
932
+ "acc": 0.41241496598639454,
933
+ "nll": 1.4535510540008545,
934
+ "brier": 0.707725465297699,
935
+ "ece": 0.025285047713388412,
936
+ "aurc": 0.48559289517381626,
937
+ "acc_at_80": 0.45483528161530284,
938
+ "acc_at_50": 0.5102040816326531,
939
+ "chance": 0.14285714285714282,
940
+ "mean_conf": 0.4194914400577545,
941
+ "sec": 13.0,
942
+ "heldout": false
943
+ },
944
+ "hate_speech_scales": {
945
+ "n": 7500,
946
+ "acc": 0.5602666666666667,
947
+ "nll": 0.9815105199813843,
948
+ "brier": 0.5425266623497009,
949
+ "ece": 0.008547326453526816,
950
+ "aurc": 0.2667196222552012,
951
+ "acc_at_80": 0.608,
952
+ "acc_at_50": 0.7093333333333334,
953
+ "chance": 0.22666666666666674,
954
+ "mean_conf": 0.5625489354133606,
955
+ "sec": 5.9,
956
+ "heldout": false,
957
+ "per_q": [
958
+ 0.706,
959
+ 0.47733333333333333,
960
+ 0.44666666666666666,
961
+ 0.5793333333333334,
962
+ 0.592
963
+ ]
964
+ },
965
+ "liar2": {
966
+ "n": 1500,
967
+ "acc": 0.37133333333333335,
968
+ "nll": 1.4664359092712402,
969
+ "brier": 0.717261552810669,
970
+ "ece": 0.010819464365641272,
971
+ "aurc": 0.47753927157183623,
972
+ "acc_at_80": 0.4066666666666667,
973
+ "acc_at_50": 0.4786666666666667,
974
+ "chance": 0.16666666666666666,
975
+ "mean_conf": 0.3717862069606781,
976
+ "sec": 3.8,
977
+ "heldout": false
978
+ },
979
+ "prosocial_safety": {
980
+ "n": 1500,
981
+ "acc": 0.5246666666666666,
982
+ "nll": 1.1752827167510986,
983
+ "brier": 0.5942376852035522,
984
+ "ece": 0.033485592842102035,
985
+ "aurc": 0.2778985096517121,
986
+ "acc_at_80": 0.5908333333333333,
987
+ "acc_at_50": 0.7,
988
+ "chance": 0.2,
989
+ "mean_conf": 0.5058215260505676,
990
+ "sec": 2.4,
991
+ "heldout": false
992
+ },
993
+ "ultrafeedback_pref": {
994
+ "n": 1500,
995
+ "acc": 0.758,
996
+ "nll": 0.49144116044044495,
997
+ "brier": 0.32426393032073975,
998
+ "ece": 0.015855272690455117,
999
+ "aurc": 0.11929765138761142,
1000
+ "acc_at_80": 0.8133333333333334,
1001
+ "acc_at_50": 0.88,
1002
+ "chance": 0.5,
1003
+ "mean_conf": 0.7586716413497925,
1004
+ "sec": 10.6,
1005
+ "heldout": false
1006
+ },
1007
+ "shp": {
1008
+ "n": 1500,
1009
+ "acc": 0.746,
1010
+ "nll": 0.517654299736023,
1011
+ "brier": 0.34469103813171387,
1012
+ "ece": 0.018088009436925262,
1013
+ "aurc": 0.13967686763081924,
1014
+ "acc_at_80": 0.7875,
1015
+ "acc_at_50": 0.8626666666666667,
1016
+ "chance": 0.5,
1017
+ "mean_conf": 0.7509077191352844,
1018
+ "sec": 7.7,
1019
+ "heldout": false
1020
+ },
1021
+ "hh_rlhf": {
1022
+ "n": 1488,
1023
+ "acc": 0.6404569892473119,
1024
+ "nll": 0.6205067038536072,
1025
+ "brier": 0.4332164227962494,
1026
+ "ece": 0.02331842826579209,
1027
+ "aurc": 0.24516950655920614,
1028
+ "acc_at_80": 0.6714285714285714,
1029
+ "acc_at_50": 0.7271505376344086,
1030
+ "chance": 0.5,
1031
+ "mean_conf": 0.6637664437294006,
1032
+ "sec": 6.1,
1033
+ "heldout": false
1034
+ },
1035
+ "arena_pref": {
1036
+ "n": 1500,
1037
+ "acc": 0.474,
1038
+ "nll": 1.1644456386566162,
1039
+ "brier": 0.6687793731689453,
1040
+ "ece": 0.14372908343871435,
1041
+ "aurc": 0.4160631567711759,
1042
+ "acc_at_80": 0.5008333333333334,
1043
+ "acc_at_50": 0.5506666666666666,
1044
+ "chance": 0.3333333333333333,
1045
+ "mean_conf": 0.6177290678024292,
1046
+ "sec": 10.0,
1047
+ "heldout": true
1048
+ },
1049
+ "reward_bench": {
1050
+ "n": 1500,
1051
+ "acc": 0.7993333333333333,
1052
+ "nll": 0.4464208483695984,
1053
+ "brier": 0.28693464398384094,
1054
+ "ece": 0.04333787786960602,
1055
+ "aurc": 0.09299298896270519,
1056
+ "acc_at_80": 0.8541666666666666,
1057
+ "acc_at_50": 0.9053333333333333,
1058
+ "chance": 0.5,
1059
+ "mean_conf": 0.7756828665733337,
1060
+ "sec": 8.0,
1061
+ "heldout": true
1062
+ },
1063
+ "glaive_tools": {
1064
+ "n": 1348,
1065
+ "acc": 0.9458456973293768,
1066
+ "nll": 0.1641840934753418,
1067
+ "brier": 0.08934415131807327,
1068
+ "ece": 0.00926415945019139,
1069
+ "aurc": 0.010922496406982686,
1070
+ "acc_at_80": 0.9768089053803339,
1071
+ "acc_at_50": 0.9955489614243324,
1072
+ "chance": 0.16666666666666666,
1073
+ "mean_conf": 0.945382297039032,
1074
+ "sec": 3.1,
1075
+ "heldout": false
1076
+ },
1077
+ "hermes_tools": {
1078
+ "n": 1500,
1079
+ "acc": 0.7373333333333333,
1080
+ "nll": 0.46730339527130127,
1081
+ "brier": 0.3166651129722595,
1082
+ "ece": 0.15524024434884393,
1083
+ "aurc": 0.05736832196708593,
1084
+ "acc_at_80": 0.8666666666666667,
1085
+ "acc_at_50": 0.9893333333333333,
1086
+ "chance": 0.125,
1087
+ "mean_conf": 0.8914145231246948,
1088
+ "sec": 9.5,
1089
+ "heldout": true
1090
+ },
1091
+ "copa": {
1092
+ "n": 100,
1093
+ "acc": 0.93,
1094
+ "nll": 0.21767663955688477,
1095
+ "brier": 0.11604493111371994,
1096
+ "ece": 0.04409654974937439,
1097
+ "aurc": 0.024093551673163044,
1098
+ "acc_at_80": 0.975,
1099
+ "acc_at_50": 0.98,
1100
+ "chance": 0.5,
1101
+ "mean_conf": 0.9102426767349243,
1102
+ "sec": 0.2,
1103
+ "heldout": false
1104
+ },
1105
+ "wic": {
1106
+ "n": 638,
1107
+ "acc": 0.713166144200627,
1108
+ "nll": 0.5775361657142639,
1109
+ "brier": 0.3854121267795563,
1110
+ "ece": 0.051907789651129334,
1111
+ "aurc": 0.20507789115696606,
1112
+ "acc_at_80": 0.7666666666666667,
1113
+ "acc_at_50": 0.8087774294670846,
1114
+ "chance": 0.5,
1115
+ "mean_conf": 0.7423800826072693,
1116
+ "sec": 1.0,
1117
+ "heldout": false
1118
+ },
1119
+ "multirc": {
1120
+ "n": 1500,
1121
+ "acc": 0.892,
1122
+ "nll": 0.28638261556625366,
1123
+ "brier": 0.1648022085428238,
1124
+ "ece": 0.025759932279586766,
1125
+ "aurc": 0.03694649471413434,
1126
+ "acc_at_80": 0.9441666666666667,
1127
+ "acc_at_50": 0.964,
1128
+ "chance": 0.5,
1129
+ "mean_conf": 0.8986733555793762,
1130
+ "sec": 7.1,
1131
+ "heldout": false
1132
+ },
1133
+ "cb": {
1134
+ "n": 56,
1135
+ "acc": 0.8928571428571429,
1136
+ "nll": 0.39255955815315247,
1137
+ "brier": 0.18121369183063507,
1138
+ "ece": 0.11490496354443687,
1139
+ "aurc": 0.02600997889991392,
1140
+ "acc_at_80": 0.9777777777777777,
1141
+ "acc_at_50": 0.9642857142857143,
1142
+ "chance": 0.3333333333333333,
1143
+ "mean_conf": 0.8255382180213928,
1144
+ "sec": 0.2,
1145
+ "heldout": true
1146
+ },
1147
+ "fever": {
1148
+ "n": 1500,
1149
+ "acc": 0.8813333333333333,
1150
+ "nll": 0.33560705184936523,
1151
+ "brier": 0.17902693152427673,
1152
+ "ece": 0.021976249794165285,
1153
+ "aurc": 0.03778744252180358,
1154
+ "acc_at_80": 0.9416666666666667,
1155
+ "acc_at_50": 0.9693333333333334,
1156
+ "chance": 0.3333333333333333,
1157
+ "mean_conf": 0.8861227035522461,
1158
+ "sec": 2.7,
1159
+ "heldout": false
1160
+ },
1161
+ "wiki_qa": {
1162
+ "n": 879,
1163
+ "acc": 0.906712172923777,
1164
+ "nll": 0.2421109974384308,
1165
+ "brier": 0.14265964925289154,
1166
+ "ece": 0.031170608783066635,
1167
+ "aurc": 0.02204155557599466,
1168
+ "acc_at_80": 0.9559032716927454,
1169
+ "acc_at_50": 0.9818181818181818,
1170
+ "chance": 0.5,
1171
+ "mean_conf": 0.8812997341156006,
1172
+ "sec": 1.5,
1173
+ "heldout": false
1174
+ },
1175
+ "msmarco_rel": {
1176
+ "n": 1446,
1177
+ "acc": 0.6708160442600276,
1178
+ "nll": 0.5872405171394348,
1179
+ "brier": 0.40445321798324585,
1180
+ "ece": 0.030953758072886067,
1181
+ "aurc": 0.2003357991113142,
1182
+ "acc_at_80": 0.7147796024200519,
1183
+ "acc_at_50": 0.7745504840940526,
1184
+ "chance": 0.5,
1185
+ "mean_conf": 0.68407142162323,
1186
+ "sec": 3.0,
1187
+ "heldout": false
1188
+ },
1189
+ "medmcqa": {
1190
+ "n": 1500,
1191
+ "acc": 0.5046666666666667,
1192
+ "nll": 1.1041309833526611,
1193
+ "brier": 0.597272515296936,
1194
+ "ece": 0.045152134160200745,
1195
+ "aurc": 0.3042279038689118,
1196
+ "acc_at_80": 0.5558333333333333,
1197
+ "acc_at_50": 0.6573333333333333,
1198
+ "chance": 0.25,
1199
+ "mean_conf": 0.5472774505615234,
1200
+ "sec": 2.4,
1201
+ "heldout": false
1202
+ },
1203
+ "quality": {
1204
+ "n": 1500,
1205
+ "acc": 0.5133333333333333,
1206
+ "nll": 1.2461479902267456,
1207
+ "brier": 0.646840512752533,
1208
+ "ece": 0.13738841543594996,
1209
+ "aurc": 0.33496979709462643,
1210
+ "acc_at_80": 0.5583333333333333,
1211
+ "acc_at_50": 0.6373333333333333,
1212
+ "chance": 0.25,
1213
+ "mean_conf": 0.6507217884063721,
1214
+ "sec": 21.8,
1215
+ "heldout": true
1216
+ },
1217
+ "xstory_cloze": {
1218
+ "n": 1500,
1219
+ "acc": 0.9646666666666667,
1220
+ "nll": 0.11040383577346802,
1221
+ "brier": 0.05835951864719391,
1222
+ "ece": 0.02761857227484382,
1223
+ "aurc": 0.004502607034060362,
1224
+ "acc_at_80": 0.9941666666666666,
1225
+ "acc_at_50": 0.9986666666666667,
1226
+ "chance": 0.5,
1227
+ "mean_conf": 0.9391650557518005,
1228
+ "sec": 2.4,
1229
+ "heldout": true
1230
+ },
1231
+ "abstain_probe": {
1232
+ "n": 1500,
1233
+ "acc": 0.7846666666666666,
1234
+ "nll": 0.6353165507316589,
1235
+ "brier": 0.2965729832649231,
1236
+ "ece": 0.03740532821416854,
1237
+ "aurc": 0.057853419152916634,
1238
+ "acc_at_80": 0.8808333333333334,
1239
+ "acc_at_50": 0.9626666666666667,
1240
+ "chance": 0.17950317460317464,
1241
+ "mean_conf": 0.8122005462646484,
1242
+ "sec": 3.5,
1243
+ "heldout": true
1244
+ },
1245
+ "toolace": {
1246
+ "n": 1000,
1247
+ "acc": 0.932,
1248
+ "nll": 0.22606024146080017,
1249
+ "brier": 0.10948573052883148,
1250
+ "ece": 0.03118352198600772,
1251
+ "aurc": 0.019889396033981634,
1252
+ "acc_at_80": 0.97125,
1253
+ "acc_at_50": 0.99,
1254
+ "chance": 0.12470555555555554,
1255
+ "mean_conf": 0.9266018867492676,
1256
+ "sec": 4.0,
1257
+ "heldout": false
1258
+ }
1259
+ },
1260
+ "agg": {
1261
+ "in_task": {
1262
+ "acc": 0.8069814634685621,
1263
+ "nll": 0.4598948841303354,
1264
+ "brier": 0.25660377455642447,
1265
+ "ece": 0.02860716135081228,
1266
+ "aurc": 0.09736104400332174,
1267
+ "acc_at_80": 0.8593853585400241,
1268
+ "chance": 0.32912234745754104
1269
+ },
1270
+ "heldout": {
1271
+ "acc": 0.7450346431863512,
1272
+ "nll": 0.6336416278196417,
1273
+ "brier": 0.34327830520013103,
1274
+ "ece": 0.0751272948477249,
1275
+ "aurc": 0.1378520558704769,
1276
+ "acc_at_80": 0.8024372456758276,
1277
+ "chance": 0.32053884112032693
1278
+ },
1279
+ "n_in": 64,
1280
+ "n_heldout": 23
1281
+ },
1282
+ "model": "runs/r3_v2/model"
1283
+ }
generation_config.json ADDED
@@ -0,0 +1,6 @@
 
 
 
 
 
 
 
1
+ {
2
+ "_from_model_config": true,
3
+ "eos_token_id": 248044,
4
+ "transformers_version": "5.17.0",
5
+ "use_cache": true
6
+ }
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:86ace0629908bff7583e67505958334f1ffd584cd9ed407e8c0fb23ef0249e0d
3
+ size 3763692048
tokenizer.json ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:06b9509352d2af50381ab2247e083b80d32d5c0aba91c272ca9ff729b6a0e523
3
+ size 19989325
tokenizer_config.json ADDED
@@ -0,0 +1,32 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "add_prefix_space": false,
3
+ "audio_bos_token": "<|audio_start|>",
4
+ "audio_eos_token": "<|audio_end|>",
5
+ "audio_token": "<|audio_pad|>",
6
+ "backend": "tokenizers",
7
+ "bos_token": null,
8
+ "clean_up_tokenization_spaces": false,
9
+ "eos_token": "<|endoftext|>",
10
+ "errors": "replace",
11
+ "image_token": "<|image_pad|>",
12
+ "is_local": false,
13
+ "local_files_only": false,
14
+ "model_max_length": 262144,
15
+ "model_specific_special_tokens": {
16
+ "audio_bos_token": "<|audio_start|>",
17
+ "audio_eos_token": "<|audio_end|>",
18
+ "audio_token": "<|audio_pad|>",
19
+ "image_token": "<|image_pad|>",
20
+ "video_token": "<|video_pad|>",
21
+ "vision_bos_token": "<|vision_start|>",
22
+ "vision_eos_token": "<|vision_end|>"
23
+ },
24
+ "pad_token": "<|endoftext|>",
25
+ "pretokenize_regex": "(?i:'s|'t|'re|'ve|'m|'ll|'d)|[^\\r\\n\\p{L}\\p{N}]?[\\p{L}\\p{M}]+|\\p{N}| ?[^\\s\\p{L}\\p{M}\\p{N}]+[\\r\\n]*|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+",
26
+ "split_special_tokens": false,
27
+ "tokenizer_class": "Qwen2Tokenizer",
28
+ "unk_token": null,
29
+ "video_token": "<|video_pad|>",
30
+ "vision_bos_token": "<|vision_start|>",
31
+ "vision_eos_token": "<|vision_end|>"
32
+ }