"""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, }