File size: 2,810 Bytes
7b875ae
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
# -*- 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]