azharmo commited on
Commit
baf5397
·
verified ·
1 Parent(s): 701f4a7

Upload jev_toy/model.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. jev_toy/model.py +198 -0
jev_toy/model.py ADDED
@@ -0,0 +1,198 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ jev_toy/model.py
3
+
4
+ A faithful-from-the-interface, toy-scale "System One" model.
5
+
6
+ WHAT THIS IS (honest framing)
7
+ -----------------------------
8
+ The real Jev model (TypeSafe AI, 2026) is proprietary and unpublished. Its
9
+ internals are NOT public. What IS public and repeatedly demonstrated:
10
+ - input = a `state` (text/JSON) + a set of TYPED `questions`
11
+ - output = per-question typed, probabilistic decisions
12
+ * choice -> probability distribution over named options
13
+ * score -> a continuous score + the distribution underneath
14
+ * noul -> a probability that a yes/no statement is true
15
+ - every question on one state is answered in PARALLEL (no token-by-token
16
+ autoregressive generation), and probabilities are claimed CALIBRATED.
17
+
18
+ What we build here is OUR OWN design to reproduce that exact interface at toy
19
+ scale. Every internal (encoder, heads, loss, calibration) is our choice, NOT a
20
+ claim about Jev's internals.
21
+
22
+ ARCHITECTURE (ours) --- built so "parallel" is real, not cosmetic
23
+ -----------------------------------------------------------------
24
+ Two-tower design driven by the proven property: one state, many questions.
25
+
26
+ state question(s)
27
+ | |
28
+ v v
29
+ StateEncoder QuestionEncoder (SHARED weights)
30
+ | |
31
+ pooled h_s pooled h_q
32
+ | |
33
+ +----- concat --+
34
+ | |
35
+ v v
36
+ typed heads (noul / choice / score)
37
+
38
+ The key fact we mirror: the state is encoded ONCE. Every question in the
39
+ request consumes that same pooled state vector, each through its own head, in
40
+ PARALLEL. Adding more questions to a request does not re-run the state encoder
41
+ --- matching the documented behaviour ("evaluate every question in parallel").
42
+
43
+ State and question are encoded with the SAME transformer weights (a shared
44
+ encoder). We feed tokens once at sampling; instr-time the two streams are just
45
+ concatenated batches through one transformer.
46
+
47
+ Heads:
48
+ noul : P(statement true) = sigmoid(logit)
49
+ choice : softmax over K logits -> distribution + confidence
50
+ score : logistic(logit)*range -> bounded real value in [lo, hi]
51
+
52
+ Calibration is a separate post-training step (calibrate.py), NOT baked into the
53
+ heads, because temperature scaling needs held-out data.
54
+
55
+ Build knobs (d_model, n_layers, vocab) are small so this trains on a laptop CPU
56
+ in minutes. Larger values + cuda device -> the same code scales to a GPU.
57
+ """
58
+
59
+ from __future__ import annotations
60
+
61
+ import math
62
+ from dataclasses import dataclass
63
+
64
+ import torch
65
+ import torch.nn as nn
66
+ import torch.nn.functional as F
67
+
68
+
69
+ @dataclass
70
+ class SystemOneConfig:
71
+ vocab_size: int = 8192
72
+ d_model: int = 128
73
+ n_layers: int = 3
74
+ n_heads: int = 4
75
+ d_ff: int = 256
76
+ max_seq_len: int = 256
77
+ pad_token_id: int = 0
78
+ num_choice_heads: int = 8
79
+ score_range: tuple[float, float] = (1.0, 5.0) # logistic maps logit -> this interval
80
+
81
+
82
+ def _rope_cache(seq_len: int, dim: int, device: torch.device):
83
+ inv = 1.0 / (10000.0 ** (torch.arange(0, dim, 2).float() / dim))
84
+ pos = torch.arange(seq_len).float().unsqueeze(1)
85
+ return (pos * inv.unsqueeze(0)).to(device)
86
+
87
+
88
+ def _apply_rope(x: torch.Tensor, cache: torch.Tensor):
89
+ B, T, H, D = x.shape
90
+ x = x.float().view(B, T, H, D // 2, 2)
91
+ x0, x1 = x.unbind(-1)
92
+ c = cache[:T].view(1, T, 1, D // 2)
93
+ cos, sin = c.cos(), c.sin()
94
+ out = torch.stack([x0 * cos - x1 * sin, x1 * cos + x0 * sin], -1).view(B, T, H, D)
95
+ return out.to(x.dtype)
96
+
97
+
98
+ class _Attn(nn.Module):
99
+ def __init__(self, cfg: SystemOneConfig):
100
+ super().__init__()
101
+ self.h = cfg.n_heads
102
+ self.dh = cfg.d_model // cfg.n_heads
103
+ self.wq = nn.Linear(cfg.d_model, cfg.d_model, bias=False)
104
+ self.wk = nn.Linear(cfg.d_model, cfg.d_model, bias=False)
105
+ self.wv = nn.Linear(cfg.d_model, cfg.d_model, bias=False)
106
+ self.wo = nn.Linear(cfg.d_model, cfg.d_model)
107
+ self.rope = None
108
+
109
+ def forward(self, x):
110
+ B, T, D = x.shape
111
+ if self.rope is None or self.rope.device != x.device:
112
+ self.rope = _rope_cache(x.shape[1], self.dh, x.device)
113
+ q = self.wq(x).view(B, T, self.h, self.dh).transpose(1, 2)
114
+ k = self.wk(x).view(B, T, self.h, self.dh).transpose(1, 2)
115
+ v = self.wv(x).view(B, T, self.h, self.dh).transpose(1, 2)
116
+ q = _apply_rope(q, self.rope)
117
+ k = _apply_rope(k, self.rope)
118
+ att = (q @ k.transpose(-2, -1)) / math.sqrt(self.dh)
119
+ # No causal mask: a decision model attends to the whole context.
120
+ h = F.softmax(att, dim=-1) @ v
121
+ return self.wo(h.transpose(1, 2).reshape(B, T, D))
122
+
123
+
124
+ class _Block(nn.Module):
125
+ def __init__(self, cfg):
126
+ super().__init__()
127
+ self.ln1 = nn.LayerNorm(cfg.d_model)
128
+ self.attn = _Attn(cfg)
129
+ self.ln2 = nn.LayerNorm(cfg.d_model)
130
+ self.ff = nn.Sequential(
131
+ nn.Linear(cfg.d_model, cfg.d_ff), nn.GELU(), nn.Linear(cfg.d_ff, cfg.d_model)
132
+ )
133
+
134
+ def forward(self, x):
135
+ x = x + self.attn(self.ln1(x))
136
+ x = x + self.ff(self.ln2(x))
137
+ return x
138
+
139
+
140
+ class _Encoder(nn.Module):
141
+ """Shared transformer encoder for state AND question text."""
142
+
143
+ def __init__(self, cfg: SystemOneConfig):
144
+ super().__init__()
145
+ self.cfg = cfg
146
+ self.embed = nn.Embedding(cfg.vocab_size, cfg.d_model, padding_idx=cfg.pad_token_id)
147
+ self.blocks = nn.ModuleList([_Block(cfg) for _ in range(cfg.n_layers)])
148
+ self.ln = nn.LayerNorm(cfg.d_model)
149
+
150
+ def forward(self, ids, mask):
151
+ x = self.embed(ids) * mask.unsqueeze(-1).float()
152
+ # position comes from RoPE inside attention
153
+ for b in self.blocks:
154
+ x = b(x)
155
+ x = self.ln(x)
156
+ pooled = (x * mask.unsqueeze(-1).float()).sum(1) / mask.sum(1, keepdim=True).clamp_min(1)
157
+ return pooled
158
+
159
+
160
+ class SystemOneModel(nn.Module):
161
+ """state -> one pooled vector; each question -> its own typed head."""
162
+
163
+ def __init__(self, cfg: SystemOneConfig):
164
+ super().__init__()
165
+ self.cfg = cfg
166
+ self.encoder = _Encoder(cfg)
167
+ # Cross-tower fusion per head.
168
+ fusion_dim = cfg.d_model * 2
169
+ self.merge = nn.Linear(fusion_dim, cfg.d_model)
170
+ self.noul_head = nn.Linear(cfg.d_model, 1)
171
+ self.score_head = nn.Linear(cfg.d_model, 1)
172
+ self.choice_head = nn.Linear(cfg.d_model, cfg.num_choice_heads)
173
+
174
+ def encode_state(self, s_ids, s_mask):
175
+ """Single forward pass over the state -> [B, D]. Called ONCE per request."""
176
+ return self.encoder(s_ids, s_mask)
177
+
178
+ def answer(self, h_state, q_ids, q_mask, q_type):
179
+ """
180
+ Answer a batch of questions off a shared state vector.
181
+ h_state: [B_state, D] (one row per state)
182
+ q_ids/q_mask: question tokens [B_q, T]; B_q may be B_state * n_questions
183
+ q_type: list/str per row: 'noul' | 'choice' | 'score'
184
+ Returns a list of per-row decision dicts (heads applied per type).
185
+ """
186
+ h_q = self.encoder(q_ids, q_mask) # [B_q, D]
187
+ # broadcast state rows to question rows
188
+ h_s = h_state # expected already expanded to B_q
189
+ h = F.gelu(self.merge(torch.cat([h_s, h_q], dim=-1)))
190
+ logits = {
191
+ "noul": self.noul_head(h).squeeze(-1),
192
+ "score": self.score_head(h).squeeze(-1),
193
+ "choice": self.choice_head(h),
194
+ }
195
+ return logits, h
196
+
197
+ def n_params(self):
198
+ return sum(p.numel() for p in self.parameters())