Upload jev_toy/serve.py with huggingface_hub
Browse files- jev_toy/serve.py +133 -0
jev_toy/serve.py
ADDED
|
@@ -0,0 +1,133 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
jev_toy/serve.py
|
| 3 |
+
|
| 4 |
+
A mirror of the System One request/response API, backed by our toy model.
|
| 5 |
+
Shows the proven interface in action:
|
| 6 |
+
|
| 7 |
+
POST: { state: "...", questions: { <name>: {type, instructions, ...} } }
|
| 8 |
+
-> { <name>: { <typed result with probabilities> } }
|
| 9 |
+
|
| 10 |
+
All questions are answered from ONE encoder pass over the state (parallel).
|
| 11 |
+
|
| 12 |
+
Usage:
|
| 13 |
+
source .venv/bin/activate
|
| 14 |
+
python -m jev_toy.serve --ckpt checkpoints/model.pt
|
| 15 |
+
"""
|
| 16 |
+
|
| 17 |
+
from __future__ import annotations
|
| 18 |
+
|
| 19 |
+
import argparse
|
| 20 |
+
|
| 21 |
+
import torch
|
| 22 |
+
import torch.nn.functional as F
|
| 23 |
+
|
| 24 |
+
from jev_toy.data import tokenize
|
| 25 |
+
from jev_toy.model import SystemOneConfig, SystemOneModel
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
def build_stoi(vocab_dict):
|
| 29 |
+
return vocab_dict
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
class ToyServer:
|
| 33 |
+
def __init__(self, ckpt, device="cpu"):
|
| 34 |
+
self.device = device
|
| 35 |
+
ck = torch.load(ckpt, map_location=device)
|
| 36 |
+
self.cfg = SystemOneConfig(**ck["config"])
|
| 37 |
+
self.stoi = ck["vocab"]
|
| 38 |
+
# __oov__ key stores the oov index; default 1
|
| 39 |
+
self.oov = ck["vocab"].get("__oov__", 1)
|
| 40 |
+
self.pad = self.cfg.pad_token_id
|
| 41 |
+
self.model = SystemOneModel(self.cfg).to(device)
|
| 42 |
+
self.model.load_state_dict(ck["state_dict"])
|
| 43 |
+
self.model.eval()
|
| 44 |
+
|
| 45 |
+
def _enc(self, text, keep=None):
|
| 46 |
+
ids, mask = tokenize(text, self.stoi, self.cfg.max_seq_len, self.oov, self.pad)
|
| 47 |
+
return (
|
| 48 |
+
torch.tensor([ids], device=self.device),
|
| 49 |
+
torch.tensor([mask], device=self.device),
|
| 50 |
+
)
|
| 51 |
+
|
| 52 |
+
@torch.no_grad()
|
| 53 |
+
def answer(self, state: str, questions: dict):
|
| 54 |
+
"""questions: {name: {type, instructions, options?}}. Returns {name: decision}."""
|
| 55 |
+
s_ids, s_mask = self._enc(state)
|
| 56 |
+
h_state = self.model.encode_state(s_ids, s_mask) # ONE pass over state
|
| 57 |
+
results = {}
|
| 58 |
+
for name, q in questions.items():
|
| 59 |
+
q_ids, q_mask = self._enc(q.get("instructions", name)) # type: ignore
|
| 60 |
+
# h_state is [1, D]; one question => one row, so pass it directly.
|
| 61 |
+
logits, _ = self.model.answer(h_state, q_ids, q_mask, q["type"])
|
| 62 |
+
t = q["type"]
|
| 63 |
+
if t == "noul":
|
| 64 |
+
p = torch.sigmoid(logits["noul"][0]).item()
|
| 65 |
+
results[name] = {"type": "noul", "noul": round(p, 4), "is_true": p >= 0.5}
|
| 66 |
+
elif t == "choice":
|
| 67 |
+
opts = q.get("options") or []
|
| 68 |
+
probs = logits["choice"][0]
|
| 69 |
+
k = len(opts)
|
| 70 |
+
if k == 0:
|
| 71 |
+
k = probs.shape[0]
|
| 72 |
+
opts = [f"option_{i}" for i in range(k)]
|
| 73 |
+
probs = probs[:k]
|
| 74 |
+
probs = F.softmax(probs / probs.sum().clamp_min(1e-9), dim=-1) # renormalize over given options
|
| 75 |
+
dist = {str(opts[i]): float(f"{probs[i].item():.4f}") for i in range(k)}
|
| 76 |
+
# confidence = margin above uniform (community formula for choice confidence)
|
| 77 |
+
u = 1.0 / k
|
| 78 |
+
conf = float((probs.max() - u) / (1 - u))
|
| 79 |
+
results[name] = {
|
| 80 |
+
"type": "choice", "distribution": dist,
|
| 81 |
+
"argmax": str(opts[int(probs.argmax())]),
|
| 82 |
+
"confidence": round(conf, 4),
|
| 83 |
+
}
|
| 84 |
+
elif t == "score":
|
| 85 |
+
lo, hi = self.cfg.score_range
|
| 86 |
+
p = (hi - lo) * torch.sigmoid(logits["score"][0]) + lo
|
| 87 |
+
results[name] = {"type": "score", "score": round(float(p), 3)}
|
| 88 |
+
else:
|
| 89 |
+
raise ValueError(t)
|
| 90 |
+
return results
|
| 91 |
+
|
| 92 |
+
|
| 93 |
+
def main():
|
| 94 |
+
ap = argparse.ArgumentParser()
|
| 95 |
+
ap.add_argument("--ckpt", default="checkpoints/model.pt")
|
| 96 |
+
args = ap.parse_args()
|
| 97 |
+
|
| 98 |
+
device = "cuda" if torch.cuda.is_available() else "cpu"
|
| 99 |
+
server = ToyServer(args.ckpt, device)
|
| 100 |
+
|
| 101 |
+
# --- demo 1: support-ticket-style multi-question (the Jev flagship use case) ---
|
| 102 |
+
demo_state = "I've been trying to connect my Stripe account for 3 days, it keeps failing with a 403. I'm losing sales and my manager is anxious. Please help ASAP."
|
| 103 |
+
demo_qs = {
|
| 104 |
+
"is_urgent": {"type": "noul", "instructions": "The message conveys urgency or time-sensitivity."},
|
| 105 |
+
"department": {
|
| 106 |
+
"type": "choice",
|
| 107 |
+
"instructions": "Which team should handle this?",
|
| 108 |
+
"options": ["billing", "technical", "sales"],
|
| 109 |
+
},
|
| 110 |
+
"frustration": {
|
| 111 |
+
"type": "score",
|
| 112 |
+
"instructions": "How frustrated does the customer appear? (0=civil, 1=angry)",
|
| 113 |
+
},
|
| 114 |
+
}
|
| 115 |
+
print("=== demo: one state, three parallel typed questions ===")
|
| 116 |
+
out = server.answer(demo_state, demo_qs)
|
| 117 |
+
for k, v in out.items():
|
| 118 |
+
print(f" {k}: {v}")
|
| 119 |
+
|
| 120 |
+
# --- demo 2: routing a news headline ---
|
| 121 |
+
print("\n=== demo: route a news article ===")
|
| 122 |
+
news = "News article: NVIDIA stock hits all-time high after beating earnings estimates, analysts raise price targets."
|
| 123 |
+
ro = server.answer(news, {
|
| 124 |
+
"topic": {
|
| 125 |
+
"type": "choice", "instructions": "Which topic does this article belong to?",
|
| 126 |
+
"options": ["World", "Sports", "Business", "Sci/Tech"],
|
| 127 |
+
}
|
| 128 |
+
})
|
| 129 |
+
print(" topic:", ro["topic"])
|
| 130 |
+
|
| 131 |
+
|
| 132 |
+
if __name__ == "__main__":
|
| 133 |
+
main()
|