azharmo commited on
Commit
8d64918
·
verified ·
1 Parent(s): 8eee705

Upload jev_toy/serve.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. jev_toy/serve.py +72 -14
jev_toy/serve.py CHANGED
@@ -51,29 +51,53 @@ class ToyServer:
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] = {
@@ -83,16 +107,47 @@ class ToyServer:
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"
@@ -128,6 +183,9 @@ def main():
128
  })
129
  print(" topic:", ro["topic"])
130
 
 
 
 
131
 
132
  if __name__ == "__main__":
133
  main()
 
51
 
52
  @torch.no_grad()
53
  def answer(self, state: str, questions: dict):
54
+ """
55
+ questions: {name: {type, instructions, options?}}. Returns {name: decision}.
56
+
57
+ PARALLEL PATH (mirrors Jev's proven claim): the state is encoded ONCE
58
+ (one encoder pass), and ALL questions are encoded in a SINGLE shared
59
+ encoder pass, then fused with the state vector and sent to the typed
60
+ heads. Nothing is generated; all answers come from one forward pass on
61
+ the state + one forward pass on the question batch.
62
+ """
63
  s_ids, s_mask = self._enc(state)
64
+ h_state = self.model.encode_state(s_ids, s_mask) # ONE pass over state
65
+
66
+ # encode ALL questions in ONE batched shared-encoder pass (parallel)
67
+ names = list(questions.keys())
68
+ q_ids_all, q_mask_all, types_all = [], [], []
69
+ for name in names:
70
+ ids, mask = self._enc(questions[name].get("instructions", name))
71
+ q_ids_all.append(ids[0]); q_mask_all.append(mask[0])
72
+ types_all.append(questions[name]["type"])
73
+ q_ids = torch.stack(q_ids_all).to(self.device) # [N, T]
74
+ q_mask = torch.stack(q_mask_all).to(self.device)
75
+ h_q = self.model.encode_questions(q_ids, q_mask) # ONE pass over all questions
76
+
77
+ # fuse state with every question row and apply typed heads
78
+ h = F.gelu(self.model.merge(torch.cat([h_state.expand(q_ids.shape[0], -1), h_q], dim=-1)))
79
+ logits = {
80
+ "noul": self.model.noul_head(h).squeeze(-1),
81
+ "score": self.model.score_head(h).squeeze(-1),
82
+ "choice": self.model.choice_head(h),
83
+ }
84
+
85
  results = {}
86
+ for i, name in enumerate(names):
87
+ t = types_all[i]
 
 
 
88
  if t == "noul":
89
+ p = torch.sigmoid(logits["noul"][i]).item()
90
  results[name] = {"type": "noul", "noul": round(p, 4), "is_true": p >= 0.5}
91
  elif t == "choice":
92
+ opts = questions[name].get("options") or []
93
+ probs = logits["choice"][i]
94
  k = len(opts)
95
  if k == 0:
96
  k = probs.shape[0]
97
  opts = [f"option_{i}" for i in range(k)]
98
  probs = probs[:k]
99
+ probs = F.softmax(probs / probs.sum().clamp_min(1e-9), dim=-1)
100
+ dist = {str(opts[j]): float(f"{probs[j].item():.4f}") for j in range(k)}
 
101
  u = 1.0 / k
102
  conf = float((probs.max() - u) / (1 - u))
103
  results[name] = {
 
107
  }
108
  elif t == "score":
109
  lo, hi = self.cfg.score_range
110
+ p = (hi - lo) * torch.sigmoid(logits["score"][i]) + lo
111
  results[name] = {"type": "score", "score": round(float(p), 3)}
112
  else:
113
  raise ValueError(t)
114
  return results
115
 
116
 
117
+ def benchmark(server, state, questions, iters=20):
118
+ """Empirically show the parallel claim: encoding the state is done ONCE,
119
+ regardless of how many questions we ask. We time the full answer() for
120
+ increasing question counts and show marginal cost per extra question.
121
+ """
122
+ import time
123
+ # build up to 12 questions (reuse the 3 real ones + synthetic nouls)
124
+ qs = list(questions.items())
125
+ while len(qs) < 12:
126
+ qs.append((f"q{len(qs)}", {
127
+ "type": "noul",
128
+ "instructions": f"Is the following claim true? Additional check {len(qs)}.",
129
+ }))
130
+ rows = []
131
+ for n in [1, 3, 6, 12]:
132
+ subset = dict(qs[:n])
133
+ server.answer(state, subset) # warmup
134
+ t0 = time.perf_counter()
135
+ for _ in range(iters):
136
+ server.answer(state, subset)
137
+ dt = (time.perf_counter() - t0) / iters * 1000
138
+ rows.append((n, dt))
139
+ print("\n=== benchmark: parallel prediction ===")
140
+ print("questions | full answer ms (CPU, single request)")
141
+ for n, ms in rows:
142
+ print(f" {n:>8} | {ms:>22.2f}")
143
+ print("\nKey: adding questions grows roughly with the question-batch size,")
144
+ print("NOT with re-reading the state. The state is encoded exactly once.")
145
+
146
+
147
  def main():
148
  ap = argparse.ArgumentParser()
149
  ap.add_argument("--ckpt", default="checkpoints/model.pt")
150
+ ap.add_argument("--bench", action="store_true", help="run parallel-encoding benchmark")
151
  args = ap.parse_args()
152
 
153
  device = "cuda" if torch.cuda.is_available() else "cpu"
 
183
  })
184
  print(" topic:", ro["topic"])
185
 
186
+ if args.bench:
187
+ benchmark(server, demo_state, demo_qs)
188
+
189
 
190
  if __name__ == "__main__":
191
  main()