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 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