vagmi commited on
Commit
5f4319c
·
verified ·
1 Parent(s): d54f613

add primitives.py

Browse files
Files changed (1) hide show
  1. primitives.py +210 -0
primitives.py ADDED
@@ -0,0 +1,210 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """
3
+ primitives.py — the one place a row becomes a prompt.
4
+
5
+ TypeSafe serves three question types against a state:
6
+
7
+ choice pick one of several named options -> probabilities + confidence
8
+ score rate against ordered levels -> expected level + legend
9
+ noul is this true? -> a single probability
10
+
11
+ All three arrive with `criteria`: a description per option, or per level. The
12
+ student is asked to read those descriptions at serve time, so it has to be
13
+ trained on them — and the teacher has to judge against the same text, or the
14
+ soft label describes a prompt nobody will ever send. That is why build_data.py,
15
+ teacher.py and jev_lite.py all render through this module and none of them
16
+ format options themselves.
17
+
18
+ Row shape (a superset of what jev_lite trains on):
19
+
20
+ {"type": "choice", "state": ..., "question": ...,
21
+ "options": ["billing", "technical"],
22
+ "criteria": {"billing": "Payments, invoicing, refunds", ...},
23
+ "label": [0.7, 0.3]}
24
+
25
+ {"type": "score", "ordered": true, "options": ["Calm", "Frustrated", "Angry"],
26
+ "label": [...]} # options ARE the levels; legend is positional
27
+
28
+ {"type": "noul", "options": ["true", "false"],
29
+ "criteria": {"true": "...", "false": "..."}, "label": [p_yes, p_no]}
30
+
31
+ `criteria` is always optional: the API allows a bare option set, and some real
32
+ tasks have no meaningful description to give (RACE's options are the answer
33
+ text itself). Training on a mix teaches the model to use criteria when they are
34
+ there and cope when they are not.
35
+ """
36
+
37
+ LETTERS = "ABCDEFGHIJKLMNOPQRSTUVWXYZ"
38
+
39
+ # noul answers in TypeSafe are keyed true/false, and the reported probability is
40
+ # the one for "true" — so index 0 must be the affirmative, always.
41
+ TRUE_FALSE = ["true", "false"]
42
+ YES_WORDS = {"yes", "true", "y"}
43
+ NO_WORDS = {"no", "false", "n"}
44
+
45
+ PREFIX = "Read the state and answer the question with one option letter.\n\n<state>\n"
46
+
47
+
48
+ def kind_of(row):
49
+ """choice | score | noul, from an explicit type or the shape of the row."""
50
+ if row.get("type") in ("choice", "score", "noul"):
51
+ return row["type"]
52
+ if row.get("ordered"):
53
+ return "score"
54
+ opts = [str(o).strip().lower() for o in row["options"]]
55
+ if len(opts) == 2 and any(o in YES_WORDS for o in opts) \
56
+ and any(o in NO_WORDS for o in opts):
57
+ return "noul"
58
+ return "choice"
59
+
60
+
61
+ def normalize(row):
62
+ """Stamp `type`, and put noul rows in true/false order with true first.
63
+
64
+ Returns a new row; the answer index and soft label are permuted with the
65
+ options, so a normalized row means exactly what the original did.
66
+ """
67
+ row = dict(row)
68
+ kind = kind_of(row)
69
+ row["type"] = kind
70
+ if kind == "score":
71
+ row["ordered"] = True
72
+ return row
73
+ if kind != "noul":
74
+ return row
75
+
76
+ opts = [str(o).strip().lower() for o in row["options"]]
77
+ yes_at = next((i for i, o in enumerate(opts) if o in YES_WORDS), None)
78
+ no_at = next((i for i, o in enumerate(opts) if o in NO_WORDS), None)
79
+ if yes_at is None or no_at is None or yes_at == no_at:
80
+ row["type"] = "choice" # not actually a yes/no pair
81
+ return row
82
+
83
+ perm = [yes_at, no_at]
84
+ criteria = row.get("criteria")
85
+ if isinstance(criteria, dict) and criteria:
86
+ row["criteria"] = {TRUE_FALSE[j]: criteria.get(row["options"][i])
87
+ for j, i in enumerate(perm)
88
+ if criteria.get(row["options"][i])}
89
+ row["options"] = list(TRUE_FALSE)
90
+ if "label" in row:
91
+ row["label"] = [row["label"][i] for i in perm]
92
+ if "answer" in row:
93
+ row["answer"] = perm.index(row["answer"])
94
+ return row
95
+
96
+
97
+ def criterion_for(row, option, position):
98
+ """The description of one option: dict lookup for choice/noul, positional for score."""
99
+ criteria = row.get("criteria")
100
+ if not criteria:
101
+ return None
102
+ if isinstance(criteria, dict):
103
+ return criteria.get(option)
104
+ if isinstance(criteria, list) and position < len(criteria):
105
+ # A score's levels and its criteria are the same ordered list; only
106
+ # index into it when the row kept them apart.
107
+ return criteria[position] if criteria[position] != option else None
108
+ return None
109
+
110
+
111
+ def option_lines(row, options=None):
112
+ """A. billing — Payments, invoicing, refunds"""
113
+ options = row["options"] if options is None else options
114
+ if len(options) > len(LETTERS):
115
+ raise ValueError(f"max {len(LETTERS)} options, got {len(options)}")
116
+ lines = []
117
+ for i, option in enumerate(options):
118
+ desc = criterion_for(row, option, i)
119
+ lines.append(f"{LETTERS[i]}. {option}" + (f" — {desc}" if desc else ""))
120
+ return "\n".join(lines)
121
+
122
+
123
+ def question_block(row, options=None):
124
+ """Everything after the state: the question, its options, their criteria.
125
+
126
+ The header tells the model which primitive it is looking at. A score is the
127
+ one case where option ORDER carries meaning, so it says so out loud.
128
+ """
129
+ header = ("Levels, lowest to highest:" if kind_of(row) == "score"
130
+ else "Options:")
131
+ return f"Question: {row['question']}\n{header}\n{option_lines(row, options)}"
132
+
133
+
134
+ def build_prompt(row, options=None):
135
+ """The full teacher-facing prompt for one row."""
136
+ return (f"{PREFIX}{row['state']}\n</state>\n\n{question_block(row, options)}\n\n"
137
+ "Reply with exactly one option letter and nothing else.\nAnswer:")
138
+
139
+
140
+ # ------------------------------------------------------------------ answers
141
+
142
+ def confidence(kind, probs, score=None):
143
+ """The model's probability that the answer it just gave is the right one.
144
+
145
+ One rule, realized per type, because each type returns a different thing:
146
+ a choice returns its argmax, so confidence is that option's probability; a
147
+ score returns an expected level, so it is the mass that rounds to the level
148
+ reported. Both are directly readable — 0.8 means wrong one time in five —
149
+ which is what makes a threshold an error budget instead of a vibe.
150
+
151
+ Chosen by measurement, not taste: over 1898 held-out rows these beat
152
+ entropy, margin, chance-corrected top and collision entropy on both AUROC
153
+ (ranking right answers above wrong ones) and ECE (the number meaning what
154
+ it says). Normalized entropy, the obvious first guess, was the worst of the
155
+ lot — it is dominated by small probabilities, so it reads a decisive
156
+ 0.85/0.08/0.07 as barely-there confidence and sends good answers to review.
157
+ """
158
+ if kind == "score":
159
+ return sum(p for i, p in enumerate(probs) if abs(i - score) <= 0.5)
160
+ return max(probs)
161
+
162
+
163
+ def answer(row, probs):
164
+ """Shape a probability distribution into the API's answer for this type."""
165
+ kind = kind_of(row)
166
+ options = row["options"]
167
+ # A noul carries no confidence field: the probability IS the answer, and
168
+ # the caller thresholds it directly.
169
+ if kind == "noul":
170
+ return {"type": "noul", "noul": round(probs[0], 4)}
171
+
172
+ if kind == "score":
173
+ expected = sum(i * p for i, p in enumerate(probs))
174
+ return {"type": "score",
175
+ "score": round(expected, 4),
176
+ "legend": {str(i): o for i, o in enumerate(options)},
177
+ "probabilities": {str(i): round(p, 4) for i, p in enumerate(probs)},
178
+ "confidence": round(confidence(kind, probs, expected), 4)}
179
+ return {"type": "choice",
180
+ "choice": options[max(range(len(probs)), key=probs.__getitem__)],
181
+ "probabilities": {o: round(p, 4) for o, p in zip(options, probs)},
182
+ "confidence": round(confidence(kind, probs), 4)}
183
+
184
+
185
+ def temper(probs, temperature):
186
+ """Flatten (T>1) or sharpen (T<1) a distribution, leaving the argmax alone.
187
+
188
+ Serving-time calibration: an adapter trained against one numeric precision
189
+ and served at another produces a distribution of the wrong sharpness. The
190
+ ranking is unaffected, so accuracy does not move — only how confident the
191
+ answer claims to be, which is the part that gates actions.
192
+ """
193
+ import math
194
+ if temperature == 1.0:
195
+ return list(probs)
196
+ logs = [math.log(max(p, 1e-12)) / temperature for p in probs]
197
+ top = max(logs)
198
+ exp = [math.exp(x - top) for x in logs]
199
+ total = sum(exp)
200
+ return [x / total for x in exp]
201
+
202
+
203
+ def normalized_entropy(probs):
204
+ """0 = certain, 1 = uniform. Divided by log(n) so option counts compare."""
205
+ import math
206
+ n = len(probs)
207
+ if n < 2:
208
+ return 0.0
209
+ h = -sum(p * math.log(p) for p in probs if p > 1e-12)
210
+ return h / math.log(n)