ajh-code's picture
Publish Jev-Bonsai-Compass runtime, adapter, card and evidence
b4b0f75 verified
Raw History Blame Contribute Delete
2.68 kB
"""Letter readout extracted from the audited local analysis/letter_readout.py."""
import math,time
LETTERS = [chr(65 + i) for i in range(26)] + [chr(97 + i) for i in range(26)]
def prompt_for(state, instructions, options):
lines = "\n".join(f"[{LETTERS[i]}] {k}: {d}" for i, (k, d) in enumerate(options))
return (
f"State:\n{state}\n\nQuestion: {instructions}\nOptions:\n{lines}\n\n"
"Answer with the letter of the best option only."
)
def letter_logprobs(choice, n):
content = (choice.get("logprobs") or {}).get("content") or []
if not content:
return {}, choice.get("message", {}).get("content"), None
top = content[0].get("top_logprobs") or []
found = {}
for item in top:
token = item.get("token") or ""
key = token if token in LETTERS else token.strip()
if key in LETTERS and (key not in found or token in LETTERS):
found[key] = item.get("logprob")
return found, content[0].get("token"), content[0].get("logprob")
def readout(client, state, instructions, options, scale=1.1, request_options=None):
body = {
"model": "qwen",
"messages": [{"role": "user", "content": prompt_for(state, instructions, options)}],
"max_tokens": 1,
"temperature": 0,
"logprobs": True,
"top_logprobs": max(40, len(options)),
"chat_template_kwargs": {"enable_thinking": False},
}
if request_options:
body.update(request_options)
t0 = time.perf_counter()
resp = client.post("/v1/chat/completions", body)
dt = time.perf_counter() - t0
choice = resp["choices"][0]
found, first_token, first_lp = letter_logprobs(choice, len(options))
raw = [found.get(LETTERS[i], -30.0) for i in range(len(options))]
missing = [LETTERS[i] for i in range(len(options)) if LETTERS[i] not in found]
scaled = [v / scale for v in raw]
m = max(scaled)
exps = [math.exp(v - m) for v in scaled]
z = sum(exps)
probs = [v / z for v in exps]
order = sorted(range(len(options)), key=lambda i: probs[i], reverse=True)
usage = resp.get("usage") or {}
timings = resp.get("timings") or {}
return {
"probs": [
{"key": options[i][0], "letter": LETTERS[i], "p": probs[i], "logprob": raw[i]}
for i in order
],
"argmax": options[order[0]][0],
"first_token": first_token,
"first_logprob": first_lp,
"missing_letters": missing,
"seconds": dt,
"prompt_tokens": usage.get("prompt_tokens"),
"cached_prompt_tokens": (usage.get("prompt_tokens_details") or {}).get("cached_tokens"),
"timings": timings,
}