arbiter-general / arbiter.py
Utiric's picture
Upload folder using huggingface_hub
eb50150 verified
Raw History Blame Contribute Delete
7.93 kB
"""ARBITER standalone inference client (copy this file, no package to install).
Mirrors the exact training/eval prompt format from System One (`s1/schema.py`,
`s1/engine.py`): <state> tags with 70/30 head-tail truncation, `Question (kind):`
labels, `(A) option — description` lines, and a single logit readout at the
`Answer: (` slot. Deviating from this format degrades scores — do not rephrase it.
Requires: torch, transformers. Weights: Qwen3-0.6B-Base fine-tune (Apache-2.0).
Usage:
from arbiter import Arbiter
arb = Arbiter("Utiric/arbiter-general") # or a local directory
out = arb.decide(
state="Package never arrived, I want a refund.",
questions=[{"id": "intent", "type": "choice",
"instructions": "What does the customer want?",
"options": {"Refund": "wants money back",
"Whereabouts": "asks where the package is"}}],
)
# {"intent": {"choice": "Refund", "probabilities": {...}, "confidence": 0.61}}
"""
import json
import string
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
LETTERS = list(string.ascii_uppercase) + list(string.ascii_lowercase) # 52 slots
KIND_TAG = {"choice": "choice", "noul": "yes/no", "score": "score"}
TRUNC_MARKER = "\n[... truncated ...]\n"
def _truncate(tok, text, max_tokens):
ids = tok.encode(text, add_special_tokens=False)
if len(ids) <= max_tokens:
return ids
head = int(max_tokens * 0.7)
tail = max_tokens - head - 8
mid = tok.encode(TRUNC_MARKER, add_special_tokens=False)
return ids[:head] + mid + ids[-tail:]
def _render_option(letter, opt, desc):
opt = str(opt).strip()
if desc:
return f"({letter}) {opt} — {str(desc).strip()}"
return f"({letter}) {opt}"
def _build_ids(tok, state, qtext, kind, options, descs, max_state_tokens):
if not isinstance(state, str):
state = json.dumps(state, ensure_ascii=False, indent=1)
bos = [tok.bos_token_id] if tok.bos_token_id is not None else []
ids = list(bos) + tok.encode("<state>\n", add_special_tokens=False)
ids += _truncate(tok, state, max_state_tokens)
ids += tok.encode("\n</state>\n", add_special_tokens=False)
lines = [f"\nQuestion ({KIND_TAG.get(kind, kind)}): {str(qtext).strip()}"]
if kind == "score":
lines.append("\nLevels:")
elif kind == "choice":
lines.append("\nOptions:")
for j, o in enumerate(options):
d = descs[j] if descs and j < len(descs) else None
lines.append("\n" + _render_option(LETTERS[j], o, d))
lines.append("\nAnswer: (")
ids.extend(tok.encode("".join(lines), add_special_tokens=False))
return ids
class Arbiter:
def __init__(self, path, device=None, temperature=None, max_state_tokens=2048):
self.device = device or ("cuda" if torch.cuda.is_available() else "cpu")
dtype = torch.bfloat16 if self.device == "cuda" else torch.float32
self.tok = AutoTokenizer.from_pretrained(path, trust_remote_code=True)
if self.tok.pad_token_id is None:
self.tok.pad_token = self.tok.eos_token or "<|endoftext|>"
self.lm = (
AutoModelForCausalLM.from_pretrained(
path, dtype=dtype, trust_remote_code=True
)
.to(self.device)
.eval()
)
inner = getattr(self.lm, "model", None)
self.backbone = getattr(inner, "language_model", inner) or self.lm
cfg = self.lm.config
self.softcap = getattr(
getattr(cfg, "text_config", cfg), "final_logit_softcapping", None
)
lids = []
for L in LETTERS:
e = self.tok.encode(L, add_special_tokens=False)
assert len(e) == 1, f"letter {L!r} is not a single token in this tokenizer"
lids.append(e[0])
self.letter_ids = torch.tensor(lids, device=self.device)
self.temperature = 1.0 if temperature is None else float(temperature)
try:
import os
cfg_path = os.path.join(path, "s1_config.json")
if os.path.isfile(cfg_path):
with open(cfg_path) as f:
self.temperature = float(
json.load(f).get("temperature", self.temperature)
)
except OSError:
pass
self.max_state_tokens = max_state_tokens
@torch.no_grad()
def _probs(self, ids, nopts):
x = torch.tensor([ids], device=self.device)
mask = torch.ones_like(x)
h = self.backbone(input_ids=x, attention_mask=mask).last_hidden_state[:, -1, :]
W = self.lm.get_output_embeddings().weight[self.letter_ids]
logits = torch.nn.functional.linear(h.to(W.dtype), W).float()
if self.softcap:
logits = torch.tanh(logits / self.softcap) * self.softcap
logits = logits / self.temperature
logits[:, nopts:] = float("-inf")
return torch.softmax(logits[0, :nopts], -1).cpu().tolist()
def decide(self, state, questions):
out = {}
for q in questions:
kind = q.get("type", "choice")
text = q.get("instructions") or q.get("question") or q.get("text") or ""
if kind == "noul":
crit = q.get("criteria") or {}
opts = ["no", "yes"]
descs = (
[crit.get("false"), crit.get("true")]
if isinstance(crit, dict)
else None
)
elif kind == "score":
levels = q.get("levels") or []
opts = [str(i) for i in range(len(levels))]
descs = list(levels)
else:
o = q.get("options") or []
if isinstance(o, dict):
opts, descs = list(o.keys()), list(o.values())
else:
opts, descs = [str(x) for x in o], None
assert len(opts) <= len(LETTERS), f"max {len(LETTERS)} options per pass"
ids = _build_ids(
self.tok, state, text, kind, opts, descs, self.max_state_tokens
)
p = self._probs(ids, len(opts))
by_opt = {o: p[i] for i, o in enumerate(opts)}
top = sorted(range(len(opts)), key=lambda i: -p[i])
conf = p[top[0]] - (p[top[1]] if len(p) > 1 else 0.0)
qid = q.get("id", "q")
if kind == "noul":
out[qid] = {"noul": p[1], "confidence": conf, "probabilities": by_opt}
elif kind == "score":
out[qid] = {
"level": str(top[0]),
"score": sum(i * v for i, v in enumerate(p)),
"confidence": conf,
"probabilities": by_opt,
}
else:
out[qid] = {
"choice": opts[top[0]],
"confidence": conf,
"probabilities": by_opt,
}
return out
if __name__ == "__main__":
import sys
arb = Arbiter(sys.argv[1] if len(sys.argv) > 1 else "Utiric/arbiter-general")
print(
json.dumps(
arb.decide(
state="Kargo 20 gündür gelmedi, iade istiyorum.",
questions=[
{
"id": "intent",
"type": "choice",
"instructions": "What does the customer want?",
"options": {
"Refund": "wants money back",
"Whereabouts": "asks where the package is",
"Cancel": "wants to cancel",
"Greeting": "just saying hello",
},
}
],
),
ensure_ascii=False,
indent=1,
)
)