Darwin-180B-RSI / handler.py
SeaWolf-AI's picture
ztc: ship early ZTC probe (AUROC 0.64) + handler
7b875ae verified
Raw History Blame Contribute Delete
2.81 kB
# -*- coding: utf-8 -*-
"""Darwin-180B-RSI handler — answer + Zero-Token Confidence (ZTC) in one JSON.
ZTC reads the final-layer hidden state of the last prompt token ONCE, before generation,
and returns the probability that the answer the model is about to produce is correct.
No extra tokens are generated and no second model is needed.
Output (one item per input):
{"answer": str, "confidence": float, "ztc_score": float, "truncated": bool}
"""
from __future__ import annotations
import os
from typing import Any, Dict, List
import numpy as np
import torch
from transformers import AutoModelForImageTextToText, AutoProcessor
class ZTC:
def __init__(self, path: str):
z = np.load(path)
self.w, self.mu, self.sd = z["w"].astype(np.float32), z["mu"].astype(np.float32), z["sd"].astype(np.float32)
self.s_mean, self.s_std = float(z["s_mean"]), float(z["s_std"])
self.A, self.B = float(z["cal_A"]), float(z["cal_B"])
def score(self, h: np.ndarray):
s = ((np.asarray(h, np.float32) - self.mu) / self.sd) @ self.w
p = 1.0 / (1.0 + np.exp(-(self.A * (s - self.s_mean) / self.s_std + self.B)))
return float(s), float(p)
class EndpointHandler:
def __init__(self, path: str = ""):
self.proc = AutoProcessor.from_pretrained(path)
self.model = AutoModelForImageTextToText.from_pretrained(path, torch_dtype="auto", device_map="auto").eval()
self.ztc = ZTC(os.path.join(path, "ztc", "ztc_probe_darwin180rsi.npz"))
@torch.no_grad()
def _one(self, prompt: str, max_new_tokens: int) -> Dict[str, Any]:
msgs = [{"role": "user", "content": [{"type": "text", "text": prompt}]}]
text = self.proc.apply_chat_template(msgs, add_generation_prompt=True, tokenize=False)
enc = self.proc(text=[text], return_tensors="pt").to(self.model.device)
# 1) ZTC: one forward pass over the prompt, final layer, last token — zero generated tokens
h = self.model(**enc, output_hidden_states=True, use_cache=False).hidden_states[-1][0, -1].float().cpu().numpy()
s, p = self.ztc.score(h)
# 2) answer
out = self.model.generate(**enc, max_new_tokens=max_new_tokens, do_sample=True, temperature=1.0, top_p=0.95, top_k=20)
gen = out[0, enc["input_ids"].shape[1]:]
answer = self.proc.decode(gen, skip_special_tokens=True)
return {"answer": answer, "confidence": round(p, 4), "ztc_score": round(s, 4), "truncated": bool(len(gen) >= max_new_tokens)}
def __call__(self, data: Dict[str, Any]) -> List[Dict[str, Any]]:
inputs = data.get("inputs")
inputs = [inputs] if isinstance(inputs, str) else inputs
mnt = int((data.get("parameters") or {}).get("max_new_tokens", 32768))
return [self._one(x, mnt) for x in inputs]