Upload folder using huggingface_hub
Browse files- .gitattributes +1 -0
- README.md +240 -0
- chat_template.jinja +154 -0
- config.json +75 -0
- decider/__init__.py +0 -0
- decider/__pycache__/__init__.cpython-312.pyc +0 -0
- decider/__pycache__/infer.cpython-312.pyc +0 -0
- decider/__pycache__/model.cpython-312.pyc +0 -0
- decider/__pycache__/prompt.cpython-312.pyc +0 -0
- decider/engine.py +124 -0
- decider/infer.py +145 -0
- decider/model.py +47 -0
- decider/prompt.py +52 -0
- eval_results.json +1283 -0
- generation_config.json +6 -0
- model.safetensors +3 -0
- tokenizer.json +3 -0
- tokenizer_config.json +32 -0
.gitattributes
CHANGED
|
@@ -33,3 +33,4 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
|
|
| 33 |
*.zip filter=lfs diff=lfs merge=lfs -text
|
| 34 |
*.zst filter=lfs diff=lfs merge=lfs -text
|
| 35 |
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
|
|
|
|
|
| 33 |
*.zip filter=lfs diff=lfs merge=lfs -text
|
| 34 |
*.zst filter=lfs diff=lfs merge=lfs -text
|
| 35 |
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
| 36 |
+
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 |
+
}
|