azharmo commited on
Commit
c46cb9f
·
verified ·
1 Parent(s): 42558dd

Upload jev_toy/serve.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. 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()