Upload jev_toy/serve.py with huggingface_hub
Browse files- 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 |
-
"""
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 55 |
s_ids, s_mask = self._enc(state)
|
| 56 |
-
h_state = self.model.encode_state(s_ids, s_mask)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 57 |
results = {}
|
| 58 |
-
for
|
| 59 |
-
|
| 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"][
|
| 65 |
results[name] = {"type": "noul", "noul": round(p, 4), "is_true": p >= 0.5}
|
| 66 |
elif t == "choice":
|
| 67 |
-
opts =
|
| 68 |
-
probs = logits["choice"][
|
| 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)
|
| 75 |
-
dist = {str(opts[
|
| 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"][
|
| 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()
|