phocinae-server 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- phocinae/__init__.py +8 -0
- phocinae/engine.py +423 -0
- phocinae/main.py +20 -0
- phocinae/model.py +159 -0
- phocinae/router.py +274 -0
- phocinae/safetensors_lite.py +54 -0
- phocinae/server.py +279 -0
- phocinae/tokenizer.py +223 -0
- phocinae_server-0.1.0.dist-info/METADATA +196 -0
- phocinae_server-0.1.0.dist-info/RECORD +13 -0
- phocinae_server-0.1.0.dist-info/WHEEL +5 -0
- phocinae_server-0.1.0.dist-info/licenses/LICENSE +201 -0
- phocinae_server-0.1.0.dist-info/top_level.txt +1 -0
phocinae/__init__.py
ADDED
|
@@ -0,0 +1,8 @@
|
|
|
1
|
+
"""phocinae-server: self-contained typed-decisions (system-one) inference server.
|
|
2
|
+
|
|
3
|
+
Zero runtime dependencies beyond fastapi + uvicorn + torch:
|
|
4
|
+
tokenizer, safetensors loading, mmBERT encoder and decision head are all
|
|
5
|
+
implemented here in pure python/torch (no laya, no transformers, no tokenizers).
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
__version__ = "0.1.0"
|
phocinae/engine.py
ADDED
|
@@ -0,0 +1,423 @@
|
|
|
1
|
+
"""Inference engine: sequence construction, batched forward, decoding.
|
|
2
|
+
|
|
3
|
+
The sequence construction mirrors laya's build_sequence protocol exactly
|
|
4
|
+
(head "[qtype] question: <ins>" + [MASK]-prefixed option markers + state +
|
|
5
|
+
EOS separators), so the server serves the released checkpoint with the same
|
|
6
|
+
input distribution it was trained with. P0 questions are mapped onto that
|
|
7
|
+
protocol:
|
|
8
|
+
|
|
9
|
+
noul -> options "false: no, the statement does not hold" /
|
|
10
|
+
"true: yes, the statement holds"; answer = P(true) >= threshold
|
|
11
|
+
choice -> options rendered as "<i>: <option text>" with the P0 list index
|
|
12
|
+
as the label; answer = index of the argmax option
|
|
13
|
+
score -> 9 ordinal levels "level 2" .. "level 10";
|
|
14
|
+
answer = 2 + argmax (an integer in 2..10)
|
|
15
|
+
|
|
16
|
+
Permutation averaging (PHOC_PERM_AVG / /v1/systemone/permute) follows the
|
|
17
|
+
canonical protocol: K=4 option orders (original, reversed, two seeded
|
|
18
|
+
shuffles, seed 0) with choice probabilities averaged before argmax.
|
|
19
|
+
"""
|
|
20
|
+
|
|
21
|
+
import json
|
|
22
|
+
import logging
|
|
23
|
+
import math
|
|
24
|
+
import os
|
|
25
|
+
import random
|
|
26
|
+
import threading
|
|
27
|
+
import time
|
|
28
|
+
|
|
29
|
+
import torch
|
|
30
|
+
|
|
31
|
+
from .model import DecisionModel, MMBertEncoder
|
|
32
|
+
from .tokenizer import META, GemmaTokenizer
|
|
33
|
+
from .safetensors_lite import load_safetensors
|
|
34
|
+
|
|
35
|
+
log = logging.getLogger("phocinae.engine")
|
|
36
|
+
|
|
37
|
+
QTYPES = {"choice": 0, "score": 1, "noul": 2}
|
|
38
|
+
|
|
39
|
+
DEFAULT_INS = {
|
|
40
|
+
"choice": "Pick the option that best matches the statement.",
|
|
41
|
+
"score": "Rate the statement.",
|
|
42
|
+
"noul": "Does the statement hold?",
|
|
43
|
+
}
|
|
44
|
+
NOUL_OPTIONS = [
|
|
45
|
+
"false: no, the statement does not hold",
|
|
46
|
+
"true: yes, the statement holds",
|
|
47
|
+
]
|
|
48
|
+
SCORE_LEVELS = 9 # answers 2..10
|
|
49
|
+
MAX_OPTIONS = 255
|
|
50
|
+
MAX_QUESTIONS = 64
|
|
51
|
+
MAX_POS = 8192 # encoder max_position_embeddings
|
|
52
|
+
PERM_K = 4 # canonical permutation averaging order count
|
|
53
|
+
BUDGET_TIERS = [(192, 512), (512, 1024), (1024, 2048), (2048, 4096),
|
|
54
|
+
(4096, MAX_POS), (MAX_POS, MAX_POS)]
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
class Engine:
|
|
58
|
+
def __init__(self, model_dir, device="auto", model_name="Phocinae-Largha-150M-v1",
|
|
59
|
+
perm_avg=False, warm=True, threads=None):
|
|
60
|
+
self.model_name = model_name
|
|
61
|
+
self.perm_avg = perm_avg
|
|
62
|
+
if device in ("auto", None, ""):
|
|
63
|
+
device = "cuda" if torch.cuda.is_available() else "cpu"
|
|
64
|
+
self.device = torch.device(device)
|
|
65
|
+
self.is_cuda = self.device.type == "cuda"
|
|
66
|
+
self.dtype = torch.float16 if self.is_cuda else torch.float32
|
|
67
|
+
if not self.is_cuda and threads:
|
|
68
|
+
torch.set_num_threads(int(threads))
|
|
69
|
+
self._lock = threading.Lock()
|
|
70
|
+
# Constant question-head texts recur on every request: cache once.
|
|
71
|
+
self._head_cache = {}
|
|
72
|
+
# Per-run encode memo for option strings (perm orders repeat them).
|
|
73
|
+
self._opt_cache = {}
|
|
74
|
+
# Lazily torch.compile'd model on CUDA (reduce-overhead, dynamic shapes);
|
|
75
|
+
# falls back to eager if compilation fails.
|
|
76
|
+
self._compiled = None
|
|
77
|
+
self._compile_failed = False
|
|
78
|
+
|
|
79
|
+
with open(os.path.join(model_dir, "encoder", "config.json"),
|
|
80
|
+
encoding="utf-8") as fh:
|
|
81
|
+
ecfg = json.load(fh)
|
|
82
|
+
with open(os.path.join(model_dir, "rl_agent_config.json"),
|
|
83
|
+
encoding="utf-8") as fh:
|
|
84
|
+
self.cfg = json.load(fh)
|
|
85
|
+
self.max_len = int(self.cfg.get("max_len", 512))
|
|
86
|
+
self.head_max_len = int(self.cfg.get("head_max_len", 192))
|
|
87
|
+
|
|
88
|
+
log.info("loading tokenizer from %s", model_dir)
|
|
89
|
+
t0 = time.time()
|
|
90
|
+
self.tok = GemmaTokenizer(os.path.join(model_dir, "tokenizer"))
|
|
91
|
+
self._verify_vocab()
|
|
92
|
+
log.info("tokenizer ready in %.2fs", time.time() - t0)
|
|
93
|
+
|
|
94
|
+
log.info("loading weights from %s", os.path.join(model_dir, "model.safetensors"))
|
|
95
|
+
t0 = time.time()
|
|
96
|
+
full_layers = {i for i, t in enumerate(ecfg["layer_types"])
|
|
97
|
+
if t == "full_attention"}
|
|
98
|
+
encoder = MMBertEncoder(
|
|
99
|
+
vocab=ecfg["vocab_size"],
|
|
100
|
+
d=ecfg["hidden_size"],
|
|
101
|
+
n_layers=ecfg["num_hidden_layers"],
|
|
102
|
+
heads=ecfg["num_attention_heads"],
|
|
103
|
+
inter=ecfg["intermediate_size"],
|
|
104
|
+
eps=float(ecfg.get("layer_norm_eps", 1e-5)),
|
|
105
|
+
full_layers=full_layers,
|
|
106
|
+
theta=float(ecfg["rope_parameters"]["full_attention"]["rope_theta"]),
|
|
107
|
+
local_window=int(ecfg["local_attention"]) // 2,
|
|
108
|
+
)
|
|
109
|
+
head = DecisionModel(
|
|
110
|
+
d=ecfg["hidden_size"],
|
|
111
|
+
head_layers=int(self.cfg.get("head_layers", 2)),
|
|
112
|
+
heads=ecfg["hidden_size"] // 64,
|
|
113
|
+
dropout=float(self.cfg.get("head_dropout", 0.1)),
|
|
114
|
+
)
|
|
115
|
+
head.encoder = encoder
|
|
116
|
+
self.model = head
|
|
117
|
+
weights = load_safetensors(os.path.join(model_dir, "model.safetensors"))
|
|
118
|
+
# The checkpoint's "temperature" tensor is a dummy (all ones); the
|
|
119
|
+
# canonical calibrated temperatures ship in rl_agent_config.json.
|
|
120
|
+
weights.pop("temperature", None)
|
|
121
|
+
self.temperature = [float(t) for t in self.cfg.get("temperature",
|
|
122
|
+
[1.0, 1.0, 1.0])]
|
|
123
|
+
if len(self.temperature) != 3:
|
|
124
|
+
raise RuntimeError("expected temperature vector of size 3")
|
|
125
|
+
weights = {k: v.to(self.dtype) for k, v in weights.items()}
|
|
126
|
+
missing, unexpected = self.model.load_state_dict(weights, strict=False)
|
|
127
|
+
if missing:
|
|
128
|
+
raise RuntimeError("weights missing for: %s" % ", ".join(sorted(missing)))
|
|
129
|
+
if unexpected:
|
|
130
|
+
raise RuntimeError("unexpected weight keys: %s" % ", ".join(sorted(unexpected)))
|
|
131
|
+
self.model = self.model.to(self.device).eval()
|
|
132
|
+
log.info("model loaded in %.2fs (%s, %d params)",
|
|
133
|
+
time.time() - t0, self.device, sum(p.numel() for p in self.model.parameters()))
|
|
134
|
+
if warm:
|
|
135
|
+
self.warmup()
|
|
136
|
+
|
|
137
|
+
# ------------------------------------------------------------ vocabulary
|
|
138
|
+
def _verify_vocab(self):
|
|
139
|
+
anchors = {0: "<pad>", 1: "<eos>", 2: "<bos>", 3: "<unk>", 4: "<mask>",
|
|
140
|
+
476: META + "a", 108: "\n", 235248: META}
|
|
141
|
+
for i, want in anchors.items():
|
|
142
|
+
got = self.tok.id2tok.get(i)
|
|
143
|
+
if got != want:
|
|
144
|
+
raise RuntimeError("vocab anchor mismatch: id %d -> %r, expected %r"
|
|
145
|
+
% (i, got, want))
|
|
146
|
+
probes = {
|
|
147
|
+
"hello": [25612],
|
|
148
|
+
"hello world": [25612, 2134],
|
|
149
|
+
"the the the": [573, 573, 573],
|
|
150
|
+
"a\nb": [476, 108, 518],
|
|
151
|
+
"a": [476],
|
|
152
|
+
" \t": [235248, 226],
|
|
153
|
+
}
|
|
154
|
+
for text, want in probes.items():
|
|
155
|
+
got = self.tok.encode(text)
|
|
156
|
+
if got != want:
|
|
157
|
+
raise RuntimeError("vocab encode mismatch for %r: got %r want %r"
|
|
158
|
+
% (text, got, want))
|
|
159
|
+
rt = "hello world, 中文测试 🚀"
|
|
160
|
+
ids = self.tok.encode(rt)
|
|
161
|
+
back = self.tok.decode(ids)
|
|
162
|
+
if back.replace(" ", "") != rt.replace(" ", ""):
|
|
163
|
+
raise RuntimeError("vocab round-trip failed: %r -> %r" % (rt, back))
|
|
164
|
+
log.info("vocabulary decode-verified (%d tokens, %d merges)",
|
|
165
|
+
len(self.tok.vocab), len(self.tok.merges))
|
|
166
|
+
|
|
167
|
+
# ---------------------------------------------------------- warmup buckets
|
|
168
|
+
def warmup(self):
|
|
169
|
+
buckets = [(1, 64, 2), (1, 128, 4), (2, 128, 4), (4, 128, 4), (4, 256, 9),
|
|
170
|
+
(8, 256, 9), (16, 256, 9), (16, 512, 16), (32, 512, 16), (64, 512, 32)]
|
|
171
|
+
log.info("warming shape buckets: %s", buckets)
|
|
172
|
+
t0 = time.time()
|
|
173
|
+
with torch.no_grad():
|
|
174
|
+
for B, L, K in buckets:
|
|
175
|
+
b = {
|
|
176
|
+
"input_ids": torch.randint(5, 2000, (B, L), dtype=torch.long),
|
|
177
|
+
"attention_mask": torch.ones((B, L), dtype=torch.long),
|
|
178
|
+
"marker_pos": torch.randint(0, L, (B, K), dtype=torch.long),
|
|
179
|
+
"marker_mask": torch.ones((B, K), dtype=torch.bool),
|
|
180
|
+
"qtype": torch.zeros((B,), dtype=torch.long),
|
|
181
|
+
}
|
|
182
|
+
self._forward(b)
|
|
183
|
+
if self.is_cuda:
|
|
184
|
+
torch.cuda.synchronize()
|
|
185
|
+
log.info("warmup done in %.2fs", time.time() - t0)
|
|
186
|
+
|
|
187
|
+
# ------------------------------------------------------------ inference
|
|
188
|
+
def _amp_ctx(self):
|
|
189
|
+
if self.is_cuda:
|
|
190
|
+
return torch.autocast(self.device.type, dtype=torch.float16)
|
|
191
|
+
return torch.autocast("cpu", enabled=False)
|
|
192
|
+
|
|
193
|
+
def _forward(self, b):
|
|
194
|
+
with torch.no_grad(), self._amp_ctx():
|
|
195
|
+
m = self.model
|
|
196
|
+
if self.is_cuda and not self._compile_failed:
|
|
197
|
+
if self._compiled is None:
|
|
198
|
+
try:
|
|
199
|
+
self._compiled = torch.compile(
|
|
200
|
+
self.model, mode="reduce-overhead", dynamic=True)
|
|
201
|
+
log.info("torch.compile active (reduce-overhead, dynamic)")
|
|
202
|
+
except Exception as exc: # pragma: no cover
|
|
203
|
+
self._compile_failed = True
|
|
204
|
+
log.warning("torch.compile unavailable: %s", exc)
|
|
205
|
+
m = self._compiled if self._compiled is not None else self.model
|
|
206
|
+
logits, act = m(
|
|
207
|
+
b["input_ids"].to(self.device),
|
|
208
|
+
b["attention_mask"].to(self.device),
|
|
209
|
+
b["marker_pos"].to(self.device),
|
|
210
|
+
b["marker_mask"].to(self.device),
|
|
211
|
+
b["qtype"].to(self.device),
|
|
212
|
+
)
|
|
213
|
+
return logits.float().cpu(), torch.softmax(act.float(), -1).cpu()
|
|
214
|
+
|
|
215
|
+
# ------------------------------------------------------- sequence building
|
|
216
|
+
def _question_spec(self, q):
|
|
217
|
+
t = q["type"]
|
|
218
|
+
ins = DEFAULT_INS[t]
|
|
219
|
+
if t == "noul":
|
|
220
|
+
return t, ins, list(NOUL_OPTIONS)
|
|
221
|
+
if t == "choice":
|
|
222
|
+
return t, ins, [str(o) for o in q["options"]]
|
|
223
|
+
return t, ins, ["level %d" % (2 + i) for i in range(SCORE_LEVELS)]
|
|
224
|
+
|
|
225
|
+
def _build_sequence(self, spec, state_ids, order, head_max_len, max_len):
|
|
226
|
+
t, ins, opts = spec
|
|
227
|
+
head_ids = self._head_cache.get((t, ins))
|
|
228
|
+
if head_ids is None:
|
|
229
|
+
head_ids = self.tok.encode("%s question: %s" % (t, ins))
|
|
230
|
+
self._head_cache[(t, ins)] = head_ids
|
|
231
|
+
opt_ids = []
|
|
232
|
+
for i in order:
|
|
233
|
+
o = opts[i].replace(self.tok.mask_token, " ")
|
|
234
|
+
ot = self._opt_cache.get(o)
|
|
235
|
+
if ot is None:
|
|
236
|
+
ot = self.tok.encode(" " + o)[:48]
|
|
237
|
+
self._opt_cache[o] = ot
|
|
238
|
+
if len(self._opt_cache) > 65536:
|
|
239
|
+
self._opt_cache.clear()
|
|
240
|
+
opt_ids.append([self.tok.mask_id] + ot)
|
|
241
|
+
opt_budget = head_max_len - sum(len(o) for o in opt_ids)
|
|
242
|
+
if opt_budget < 16:
|
|
243
|
+
per = max(4, (head_max_len - 16) // max(1, len(opt_ids)))
|
|
244
|
+
opt_ids = [o[:per] for o in opt_ids]
|
|
245
|
+
opt_budget = head_max_len - sum(len(o) for o in opt_ids)
|
|
246
|
+
head_ids = head_ids[:max(8, opt_budget)]
|
|
247
|
+
ids = [self.tok.cls_token_id] + head_ids + [self.tok.sep_token_id]
|
|
248
|
+
markers = []
|
|
249
|
+
for o in opt_ids:
|
|
250
|
+
markers.append(len(ids))
|
|
251
|
+
ids.extend(o)
|
|
252
|
+
ids.append(self.tok.sep_token_id)
|
|
253
|
+
room = max(0, max_len - len(ids) - 1)
|
|
254
|
+
ids = ids + state_ids[:room] + [self.tok.sep_token_id]
|
|
255
|
+
return ids[:max_len], [m for m in markers if m < max_len]
|
|
256
|
+
|
|
257
|
+
def _build_sequence_fit(self, spec, state_ids, order=None):
|
|
258
|
+
k = len(spec[2])
|
|
259
|
+
if order is None:
|
|
260
|
+
order = list(range(k))
|
|
261
|
+
for hml, ml in BUDGET_TIERS:
|
|
262
|
+
seq, markers = self._build_sequence(spec, state_ids, order, hml, ml)
|
|
263
|
+
if len(markers) == k:
|
|
264
|
+
return seq, markers
|
|
265
|
+
raise ValueError("options exceed the %d-token model budget" % MAX_POS)
|
|
266
|
+
|
|
267
|
+
def _encode_request(self, state, questions):
|
|
268
|
+
state_ids = self.tok.encode(state.replace(self.tok.mask_token, " "))
|
|
269
|
+
items = []
|
|
270
|
+
for q in questions:
|
|
271
|
+
spec = self._question_spec(q)
|
|
272
|
+
seq, markers = self._build_sequence_fit(spec, state_ids)
|
|
273
|
+
items.append({
|
|
274
|
+
"ids": seq,
|
|
275
|
+
"markers": markers,
|
|
276
|
+
"qtype": QTYPES[q["type"]],
|
|
277
|
+
"n_opts": len(spec[2]),
|
|
278
|
+
"spec": spec,
|
|
279
|
+
"q": q,
|
|
280
|
+
})
|
|
281
|
+
return state_ids, items
|
|
282
|
+
|
|
283
|
+
# ------------------------------------------------------------- collation
|
|
284
|
+
def _collate(self, items):
|
|
285
|
+
n = len(items)
|
|
286
|
+
L = max(len(it["ids"]) for it in items)
|
|
287
|
+
K = max(len(it["markers"]) for it in items)
|
|
288
|
+
ids = [[self.tok.pad_id] * L for _ in range(n)]
|
|
289
|
+
att = [[0] * L for _ in range(n)]
|
|
290
|
+
mpos = [[0] * K for _ in range(n)]
|
|
291
|
+
mmask = [[False] * K for _ in range(n)]
|
|
292
|
+
for i, it in enumerate(items):
|
|
293
|
+
ids[i][:len(it["ids"])] = it["ids"]
|
|
294
|
+
att[i][:len(it["ids"])] = [1] * len(it["ids"])
|
|
295
|
+
for j, m in enumerate(it["markers"]):
|
|
296
|
+
mpos[i][j] = m
|
|
297
|
+
mmask[i][j] = True
|
|
298
|
+
return {
|
|
299
|
+
"input_ids": torch.tensor(ids, dtype=torch.long),
|
|
300
|
+
"attention_mask": torch.tensor(att, dtype=torch.long),
|
|
301
|
+
"marker_pos": torch.tensor(mpos, dtype=torch.long),
|
|
302
|
+
"marker_mask": torch.tensor(mmask, dtype=torch.bool),
|
|
303
|
+
"qtype": torch.tensor([it["qtype"] for it in items], dtype=torch.long),
|
|
304
|
+
}
|
|
305
|
+
|
|
306
|
+
# --------------------------------------------------------------- decoding
|
|
307
|
+
@staticmethod
|
|
308
|
+
def _probs(logits_row, k, t_scale):
|
|
309
|
+
z = [float(v) / t_scale for v in logits_row[:k]]
|
|
310
|
+
m = max(z)
|
|
311
|
+
p = [math.exp(v - m) for v in z]
|
|
312
|
+
s = sum(p)
|
|
313
|
+
return [v / s for v in p]
|
|
314
|
+
|
|
315
|
+
def _decode_rows(self, logits, act, items):
|
|
316
|
+
"""Returns per-item (answer_value, confidence, act_probability).
|
|
317
|
+
|
|
318
|
+
confidence follows laya's canonical answer_confidence:
|
|
319
|
+
float(clip(max(p[:k]), 0, 1)) — the calibrated top probability.
|
|
320
|
+
"""
|
|
321
|
+
out = []
|
|
322
|
+
for j, it in enumerate(items):
|
|
323
|
+
k = it["n_opts"]
|
|
324
|
+
t = it["q"]["type"]
|
|
325
|
+
t_scale = self.temperature[QTYPES[t]]
|
|
326
|
+
p = self._probs(logits[j], k, t_scale)
|
|
327
|
+
conf = round(float(min(1.0, max(p[:k]))), 4)
|
|
328
|
+
act_prob = round(float(act[j, 0]), 4)
|
|
329
|
+
if t == "choice":
|
|
330
|
+
ans = max(range(k), key=lambda i: p[i])
|
|
331
|
+
elif t == "score":
|
|
332
|
+
ans = 2 + max(range(k), key=lambda i: p[i])
|
|
333
|
+
else:
|
|
334
|
+
thr = it["q"].get("threshold")
|
|
335
|
+
thr = 0.5 if thr is None else float(thr)
|
|
336
|
+
ans = bool(p[1] >= thr)
|
|
337
|
+
out.append((ans, conf, act_prob, p))
|
|
338
|
+
return out
|
|
339
|
+
|
|
340
|
+
# ------------------------------------------------------------- permutation
|
|
341
|
+
def _perm_orders(self, item, rng):
|
|
342
|
+
k = item["n_opts"]
|
|
343
|
+
if item["q"]["type"] != "choice" or k < 2:
|
|
344
|
+
return [list(range(k))]
|
|
345
|
+
o1 = list(reversed(range(k)))
|
|
346
|
+
o2 = list(range(k))
|
|
347
|
+
rng.shuffle(o2)
|
|
348
|
+
o3 = list(range(k))
|
|
349
|
+
rng.shuffle(o3)
|
|
350
|
+
return [list(range(k)), o1, o2, o3]
|
|
351
|
+
|
|
352
|
+
# ----------------------------------------------------------------- run
|
|
353
|
+
def run(self, state, questions, perm=None, usage=True, with_scores=False):
|
|
354
|
+
"""Run one system-one request. Returns (answers, confidence, action, usage).
|
|
355
|
+
|
|
356
|
+
perm: None -> engine default; False -> single pass; True -> permute-avg.
|
|
357
|
+
with_scores: True -> also return the per-option calibrated probability
|
|
358
|
+
vectors as a 5th element: {qid: [p0, p1, ...]} in canonical option
|
|
359
|
+
order (choice: request option order; noul: [false, true];
|
|
360
|
+
score: levels 2..10). Choice vectors are averaged over permuted
|
|
361
|
+
orders when perm=True.
|
|
362
|
+
"""
|
|
363
|
+
if perm is None:
|
|
364
|
+
perm = self.perm_avg
|
|
365
|
+
with self._lock:
|
|
366
|
+
state_ids, items = self._encode_request(state, questions)
|
|
367
|
+
rng = random.Random(0)
|
|
368
|
+
rows = []
|
|
369
|
+
row_meta = [] # (question_idx, order)
|
|
370
|
+
for qi, it in enumerate(items):
|
|
371
|
+
orders = self._perm_orders(it, rng) if perm \
|
|
372
|
+
else [list(range(it["n_opts"]))]
|
|
373
|
+
for order in orders:
|
|
374
|
+
seq, markers = self._build_sequence_fit(it["spec"], state_ids, order)
|
|
375
|
+
rows.append({"ids": seq, "markers": markers,
|
|
376
|
+
"qtype": it["qtype"], "n_opts": it["n_opts"],
|
|
377
|
+
"spec": it["spec"], "q": it["q"]})
|
|
378
|
+
row_meta.append((qi, order))
|
|
379
|
+
b = self._collate(rows)
|
|
380
|
+
logits, act = self._forward(b)
|
|
381
|
+
decoded = self._decode_rows(logits.numpy(), act.numpy(), rows)
|
|
382
|
+
|
|
383
|
+
# regroup per question; average choice probabilities over orders
|
|
384
|
+
n_q = len(questions)
|
|
385
|
+
per_q = [[] for _ in range(n_q)]
|
|
386
|
+
for (qi, order), d in zip(row_meta, decoded):
|
|
387
|
+
per_q[qi].append((order,) + d)
|
|
388
|
+
answers, confs, actions = {}, {}, {}
|
|
389
|
+
scores = {}
|
|
390
|
+
for qi, rowsq in enumerate(per_q):
|
|
391
|
+
q = questions[qi]
|
|
392
|
+
order0 = rowsq[0][0]
|
|
393
|
+
n_opts = len(order0)
|
|
394
|
+
if len(rowsq) == 1:
|
|
395
|
+
_, ans, conf, actp, p = rowsq[0]
|
|
396
|
+
else:
|
|
397
|
+
pal = [0.0] * n_opts
|
|
398
|
+
for order, _ans, _conf, _actp, p in rowsq:
|
|
399
|
+
for pos, opt_idx in enumerate(order):
|
|
400
|
+
pal[opt_idx] += p[pos]
|
|
401
|
+
m = len(rowsq)
|
|
402
|
+
pal = [v / m for v in pal]
|
|
403
|
+
p = pal
|
|
404
|
+
conf = round(float(min(1.0, max(pal[:n_opts]))), 4)
|
|
405
|
+
actp = rowsq[0][3]
|
|
406
|
+
ans = max(range(n_opts), key=lambda i: pal[i])
|
|
407
|
+
if q["type"] == "choice":
|
|
408
|
+
pass # ans is already the option index
|
|
409
|
+
elif q["type"] == "score":
|
|
410
|
+
ans = 2 + max(range(len(p)), key=lambda i: p[i])
|
|
411
|
+
else:
|
|
412
|
+
thr = q.get("threshold")
|
|
413
|
+
thr = 0.5 if thr is None else float(thr)
|
|
414
|
+
ans = bool(p[1] >= thr)
|
|
415
|
+
answers[q["id"]] = ans
|
|
416
|
+
confs[q["id"]] = conf
|
|
417
|
+
actions[q["id"]] = {"act_probability": actp}
|
|
418
|
+
scores[q["id"]] = [round(float(v), 4) for v in p]
|
|
419
|
+
total_tokens = int(b["attention_mask"].sum().item())
|
|
420
|
+
usage = {"input_tokens": total_tokens, "output_tokens": 0}
|
|
421
|
+
if with_scores:
|
|
422
|
+
return answers, confs, actions, usage, scores
|
|
423
|
+
return answers, confs, actions, usage
|
phocinae/main.py
ADDED
|
@@ -0,0 +1,20 @@
|
|
|
1
|
+
"""uvicorn entrypoint: python -m phocinae.main"""
|
|
2
|
+
|
|
3
|
+
import logging
|
|
4
|
+
import os
|
|
5
|
+
|
|
6
|
+
import uvicorn
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
def main():
|
|
10
|
+
logging.basicConfig(
|
|
11
|
+
level=os.environ.get("PHOC_LOG", "INFO").upper(),
|
|
12
|
+
format="%(asctime)s %(levelname)s %(name)s: %(message)s")
|
|
13
|
+
host = os.environ.get("PHOC_HOST", "127.0.0.1")
|
|
14
|
+
port = int(os.environ.get("PHOC_PORT", "8155"))
|
|
15
|
+
uvicorn.run("phocinae.server:app", host=host, port=port, workers=1,
|
|
16
|
+
log_level="warning", access_log=os.environ.get("PHOC_ACCESS_LOG", "0") == "1")
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
if __name__ == "__main__":
|
|
20
|
+
main()
|
phocinae/model.py
ADDED
|
@@ -0,0 +1,159 @@
|
|
|
1
|
+
"""Pure-torch mmBERT-small encoder + typed-decisions head.
|
|
2
|
+
|
|
3
|
+
Re-implements (inference only, no runtime dependency on laya/transformers):
|
|
4
|
+
- mmBERT-small encoder: 384 hidden / 22 layers / 6 heads / GEGLU MLP /
|
|
5
|
+
RoPE (Gemma2-style, theta=160000) / sliding-window attention on all but
|
|
6
|
+
every third layer (full attention) / pre-norm, LayerNorm(eps=1e-5, no bias),
|
|
7
|
+
layer 0 attention norm is Identity.
|
|
8
|
+
- decision head: type embedding + 2-layer pre-norm TransformerEncoder
|
|
9
|
+
(nn.TransformerEncoderLayer, d=384, 6 heads, ff=1536, dropout 0.1) +
|
|
10
|
+
marker-gathered scorer + act head (d+4 -> 256 -> 2).
|
|
11
|
+
|
|
12
|
+
State-dict keys match the released checkpoint exactly
|
|
13
|
+
(encoder.layers.i.attn.Wqkv/Wo, mlp.Wi/Wo, attn_norm/mlp_norm, head.layers.*,
|
|
14
|
+
scorer.*, act_head.*, type_emb.weight).
|
|
15
|
+
"""
|
|
16
|
+
|
|
17
|
+
import math
|
|
18
|
+
|
|
19
|
+
import torch
|
|
20
|
+
import torch.nn as nn
|
|
21
|
+
import torch.nn.functional as F
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
def _rotate_half(x):
|
|
25
|
+
d = x.shape[-1] // 2
|
|
26
|
+
return torch.cat([-x[..., d:], x[..., :d]], dim=-1)
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
class _Attn(nn.Module):
|
|
30
|
+
"""Fused Wqkv attention with RoPE + SDPA (checkpoint keys attn.Wqkv/attn.Wo)."""
|
|
31
|
+
|
|
32
|
+
def __init__(self, d, heads):
|
|
33
|
+
super().__init__()
|
|
34
|
+
self.heads = heads
|
|
35
|
+
self.hd = d // heads
|
|
36
|
+
self.Wqkv = nn.Linear(d, 3 * d, bias=False)
|
|
37
|
+
self.Wo = nn.Linear(d, d, bias=False)
|
|
38
|
+
|
|
39
|
+
def forward(self, x, mask, cos, sin):
|
|
40
|
+
B, L, _ = x.shape
|
|
41
|
+
qkv = self.Wqkv(x).view(B, L, 3, self.heads, self.hd)
|
|
42
|
+
q, k, v = qkv.unbind(2)
|
|
43
|
+
q = q.transpose(1, 2)
|
|
44
|
+
k = k.transpose(1, 2)
|
|
45
|
+
v = v.transpose(1, 2)
|
|
46
|
+
q = (q.float() * cos + _rotate_half(q.float()) * sin).to(q.dtype)
|
|
47
|
+
k = (k.float() * cos + _rotate_half(k.float()) * sin).to(k.dtype)
|
|
48
|
+
o = F.scaled_dot_product_attention(
|
|
49
|
+
q, k, v, attn_mask=mask, dropout_p=0.0, is_causal=False,
|
|
50
|
+
scale=1.0 / math.sqrt(self.hd))
|
|
51
|
+
return self.Wo(o.transpose(1, 2).reshape(B, L, -1))
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
class _MLP(nn.Module):
|
|
55
|
+
"""GEGLU MLP (checkpoint keys mlp.Wi/mlp.Wo): Wo(gelu(input) * gate)."""
|
|
56
|
+
|
|
57
|
+
def __init__(self, d, inter):
|
|
58
|
+
super().__init__()
|
|
59
|
+
self.Wi = nn.Linear(d, 2 * inter, bias=False)
|
|
60
|
+
self.Wo = nn.Linear(inter, d, bias=False)
|
|
61
|
+
|
|
62
|
+
def forward(self, x):
|
|
63
|
+
inp, gate = self.Wi(x).chunk(2, dim=-1)
|
|
64
|
+
return self.Wo(F.gelu(inp) * gate)
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
class _EncoderLayer(nn.Module):
|
|
68
|
+
def __init__(self, d, heads, inter, eps, full, first):
|
|
69
|
+
super().__init__()
|
|
70
|
+
self.full = full
|
|
71
|
+
self.attn_norm = nn.Identity() if first else nn.LayerNorm(d, eps=eps, bias=False)
|
|
72
|
+
self.attn = _Attn(d, heads)
|
|
73
|
+
self.mlp_norm = nn.LayerNorm(d, eps=eps, bias=False)
|
|
74
|
+
self.mlp = _MLP(d, inter)
|
|
75
|
+
|
|
76
|
+
def forward(self, x, mask, cos, sin):
|
|
77
|
+
x = x + self.attn(self.attn_norm(x), mask, cos, sin)
|
|
78
|
+
x = x + self.mlp(self.mlp_norm(x))
|
|
79
|
+
return x
|
|
80
|
+
|
|
81
|
+
|
|
82
|
+
class MMBertEncoder(nn.Module):
|
|
83
|
+
def __init__(self, vocab, d, n_layers, heads, inter, eps, full_layers,
|
|
84
|
+
theta, local_window):
|
|
85
|
+
super().__init__()
|
|
86
|
+
self.hd = d // heads
|
|
87
|
+
self.local_window = local_window
|
|
88
|
+
self.embeddings = nn.Module()
|
|
89
|
+
self.embeddings.tok_embeddings = nn.Embedding(vocab, d, padding_idx=0)
|
|
90
|
+
self.embeddings.norm = nn.LayerNorm(d, eps=eps, bias=False)
|
|
91
|
+
self.layers = nn.ModuleList([
|
|
92
|
+
_EncoderLayer(d, heads, inter, eps, full=(i in full_layers), first=(i == 0))
|
|
93
|
+
for i in range(n_layers)
|
|
94
|
+
])
|
|
95
|
+
self.final_norm = nn.LayerNorm(d, eps=eps, bias=False)
|
|
96
|
+
inv = 1.0 / (theta ** (torch.arange(0, self.hd, 2, dtype=torch.float32) / self.hd))
|
|
97
|
+
self.register_buffer("inv_freq", inv, persistent=False)
|
|
98
|
+
|
|
99
|
+
def forward(self, input_ids, attention_mask):
|
|
100
|
+
B, L = input_ids.shape
|
|
101
|
+
dev = input_ids.device
|
|
102
|
+
x = self.embeddings.tok_embeddings(input_ids)
|
|
103
|
+
x = self.embeddings.norm(x)
|
|
104
|
+
pos = torch.arange(L, device=dev, dtype=torch.float32)
|
|
105
|
+
freqs = torch.outer(pos, self.inv_freq.to(dev)) # [L, hd/2] fp32
|
|
106
|
+
emb = torch.cat([freqs, freqs], dim=-1) # [L, hd]
|
|
107
|
+
cos = emb.cos()[None, None] # [1, 1, L, hd]
|
|
108
|
+
sin = emb.sin()[None, None]
|
|
109
|
+
pad = attention_mask.bool() # [B, L]
|
|
110
|
+
full_mask = pad[:, None, None, :] # [B, 1, 1, L]
|
|
111
|
+
band = (pos[:, None] - pos[None, :]).abs() <= self.local_window
|
|
112
|
+
slide_mask = band[None, None] & pad[:, None, None, :] # [B, 1, L, L]
|
|
113
|
+
for layer in self.layers:
|
|
114
|
+
x = layer(x, full_mask if layer.full else slide_mask, cos, sin)
|
|
115
|
+
return self.final_norm(x)
|
|
116
|
+
|
|
117
|
+
|
|
118
|
+
class DecisionModel(nn.Module):
|
|
119
|
+
"""Full decision model; checkpoint keys live at the top level:
|
|
120
|
+
|
|
121
|
+
encoder.*, head.layers.*, scorer.*, act_head.*, type_emb.weight
|
|
122
|
+
"""
|
|
123
|
+
|
|
124
|
+
def __init__(self, d, head_layers, heads, dropout):
|
|
125
|
+
super().__init__()
|
|
126
|
+
self.type_emb = nn.Embedding(3, d)
|
|
127
|
+
layer = nn.TransformerEncoderLayer(
|
|
128
|
+
d, heads, 4 * d, dropout, batch_first=True, norm_first=True)
|
|
129
|
+
self.head = nn.TransformerEncoder(
|
|
130
|
+
layer, head_layers, enable_nested_tensor=False)
|
|
131
|
+
self.scorer = nn.Sequential(
|
|
132
|
+
nn.LayerNorm(d), nn.Linear(d, d), nn.GELU(), nn.Linear(d, 1))
|
|
133
|
+
self.act_head = nn.Sequential(
|
|
134
|
+
nn.Linear(d + 4, 256), nn.GELU(), nn.Linear(256, 2))
|
|
135
|
+
|
|
136
|
+
def forward(self, input_ids, attention_mask, marker_pos, marker_mask, qtype):
|
|
137
|
+
h = self.encoder(input_ids, attention_mask)
|
|
138
|
+
h = h + self.type_emb(qtype)[:, None, :]
|
|
139
|
+
pad = ~attention_mask.bool()
|
|
140
|
+
h = self.head(h, src_key_padding_mask=pad)
|
|
141
|
+
return self._score(h, marker_pos, marker_mask)
|
|
142
|
+
|
|
143
|
+
def _score(self, h, marker_pos, marker_mask):
|
|
144
|
+
idx = marker_pos.clamp(min=0)[:, :, None].expand(-1, -1, h.size(-1))
|
|
145
|
+
m = torch.gather(h, 1, idx)
|
|
146
|
+
logits = self.scorer(m).squeeze(-1).float()
|
|
147
|
+
logits = logits.masked_fill(~marker_mask, -1e4)
|
|
148
|
+
p = torch.softmax(logits.detach(), -1)
|
|
149
|
+
k = marker_mask.sum(-1).clamp(min=2).float()
|
|
150
|
+
ent = -(p * torch.log(p.clamp_min(1e-9))).sum(-1) / torch.log(k)
|
|
151
|
+
if p.size(-1) >= 2:
|
|
152
|
+
top2 = p.topk(2, -1).values
|
|
153
|
+
else:
|
|
154
|
+
top1 = p.topk(1, -1).values
|
|
155
|
+
top2 = torch.cat([top1, torch.zeros_like(top1)], dim=-1)
|
|
156
|
+
feats = torch.stack(
|
|
157
|
+
[top2[:, 0], top2[:, 0] - top2[:, 1], ent, k / 255.0], -1)
|
|
158
|
+
act_logits = self.act_head(torch.cat([h[:, 0].float(), feats], -1))
|
|
159
|
+
return logits, act_logits
|