nx_model_spec.nx source
↩ module page · 337 lines · 12091 B
1// nx_model_spec.nx -- typed model configuration envelope.
2//
3// L4 foundational data brick. Every model-running primitive composes
4// against this: nx_gguf loader populates an NxModelSpec from
5// metadata; the inference scaffold reads it to size scratch buffers
6// and dispatch the right norm + activation + position-encoding
7// kernels.
8//
9// Per the bits-up cardinal: model-specific constants (n_heads,
10// hidden_dim, etc.) belong in a TYPED struct, not scattered as
11// magic numbers across call sites. Per the data-driven cardinal:
12// the same substrate runs Llama / Mistral / Qwen / GPT-2 / BERT
13// just by swapping NxModelSpec values.
14//
15// ===== Sealed enums ==============================================
16//
17// Norm kind:
18// NX_NORM_RMS -- Zhang/Sennrich 2019 (Llama / Mistral / Qwen)
19// NX_NORM_LAYER -- Ba/Kiros/Hinton 2016 (GPT-2 / BERT / T5 / ViT)
20// NX_NORM_NONE -- bypass (debug / experimentation)
21//
22// Activation kind (FFN gate function):
23// NX_ACT_SILU -- Ramachandran 2017 (Llama / Mistral SwiGLU)
24// NX_ACT_GELU -- Hendrycks/Gimpel 2016 (GPT-2 / BERT GeGLU)
25// NX_ACT_RELU -- Fukushima 1969 / Glorot 2011 (older models)
26//
27// Position encoding kind:
28// NX_POS_ROPE -- Su 2021 (Llama / Mistral / Qwen)
29// NX_POS_SINUSOIDAL -- Vaswani 2017 (original transformer)
30// NX_POS_LEARNED -- BERT / GPT-2 absolute-learned table
31// NX_POS_NONE -- bypass
32//
33// FFN gating kind:
34// NX_FFN_SWIGLU -- SiLU(gate) * up (Llama / Mistral)
35// NX_FFN_GEGLU -- GELU(gate) * up (newer encoder models)
36// NX_FFN_PLAIN -- act(up) (original transformer)
37//
38// genealogy_id: yaml_huggingface_config_2022 + onnx_model_metadata +
39// ggml_gguf_metadata_2024
40// lineage_id: substrate_model_spec_v1
41
42// nx_safety_envelope:
43// intended_use: AUTO_APPLIED -- primitive-specific tuning queued
44// sil_target: SIL1
45// evidence: [bulk_applied_2026-05-16, see-file-comment-for-detail]
46// verdict: NOT_YET_EVALUATED
47
48import "nx_syscalls.nx"
49import "nx_tier.nx"
50import "nx_loop.nx"
51
52// ===== Sealed-enum: NormKind ======================================
53
54const NX_NORM_RMS: nx_int = 0
55const NX_NORM_LAYER: nx_int = 1
56const NX_NORM_NONE: nx_int = 2
57const NX_NORM_N: nx_int = 3
58
59func nx_norm_is_valid(k: nx_int) -> nx_int {
60 if k < 0 { return 0 }
61 if k >= NX_NORM_N { return 0 }
62 return 1
63}
64
65// ===== Sealed-enum: ActKind =======================================
66
67const NX_ACT_SILU: nx_int = 0
68const NX_ACT_GELU: nx_int = 1
69const NX_ACT_RELU: nx_int = 2
70const NX_ACT_NONE: nx_int = 3
71const NX_ACT_N: nx_int = 4
72
73func nx_act_is_valid(k: nx_int) -> nx_int {
74 if k < 0 { return 0 }
75 if k >= NX_ACT_N { return 0 }
76 return 1
77}
78
79// ===== Sealed-enum: PosKind =======================================
80
81const NX_POS_ROPE: nx_int = 0
82const NX_POS_SINUSOIDAL: nx_int = 1
83const NX_POS_LEARNED: nx_int = 2
84const NX_POS_NONE: nx_int = 3
85const NX_POS_N: nx_int = 4
86
87func nx_pos_is_valid(k: nx_int) -> nx_int {
88 if k < 0 { return 0 }
89 if k >= NX_POS_N { return 0 }
90 return 1
91}
92
93// ===== Sealed-enum: FfnKind =======================================
94
95const NX_FFN_SWIGLU: nx_int = 0
96const NX_FFN_GEGLU: nx_int = 1
97const NX_FFN_PLAIN: nx_int = 2
98const NX_FFN_N: nx_int = 3
99
100func nx_ffn_is_valid(k: nx_int) -> nx_int {
101 if k < 0 { return 0 }
102 if k >= NX_FFN_N { return 0 }
103 return 1
104}
105
106// ===== Sealed-enum: ModelSpecVerdict ==============================
107
108const NX_MS_OK: nx_int = 0
109const NX_MS_ERR_BAD_DIM: nx_int = 1
110const NX_MS_ERR_BAD_HEAD_GEOM: nx_int = 2
111const NX_MS_ERR_BAD_KIND: nx_int = 3
112const NX_MS_ERR_BAD_VOCAB: nx_int = 4
113const NX_MS_N_VERDICTS: nx_int = 5
114
115func nx_ms_verdict_is_valid(v: nx_int) -> nx_int {
116 if v < 0 { return 0 }
117 if v >= NX_MS_N_VERDICTS { return 0 }
118 return 1
119}
120
121// ===== The struct =================================================
122
123struct NxModelSpec {
124 n_layers: nx_int,
125 hidden_dim: nx_int, // = n_heads * head_dim
126 n_heads: nx_int,
127 head_dim: nx_int,
128 n_kv_heads: nx_int, // for grouped-query attention; == n_heads if no GQA
129 ffn_dim: nx_int, // intermediate FFN dimension
130 vocab_size: nx_int,
131 max_seq_len: nx_int,
132 norm_kind: nx_int, // NX_NORM_*
133 act_kind: nx_int, // NX_ACT_*
134 pos_kind: nx_int, // NX_POS_*
135 ffn_kind: nx_int // NX_FFN_*
136}
137
138const NX_MS_STRUCT_BYTES: nx_int = 96 // 12 fields * 8
139
140func nx_model_spec_new() -> *NxModelSpec {
141 let s: *NxModelSpec = sys_mmap(NX_MS_STRUCT_BYTES) as *NxModelSpec
142 // Defaults match Llama-7B-shape values; caller overrides.
143 s.n_layers = 32
144 s.hidden_dim = 4096
145 s.n_heads = 32
146 s.head_dim = 128
147 s.n_kv_heads = 32
148 s.ffn_dim = 11008
149 s.vocab_size = 32000
150 s.max_seq_len = 2048
151 s.norm_kind = NX_NORM_RMS
152 s.act_kind = NX_ACT_SILU
153 s.pos_kind = NX_POS_ROPE
154 s.ffn_kind = NX_FFN_SWIGLU
155 return s
156}
157
158// ===== Validation =================================================
159//
160// Returns NX_MS_OK if the spec is internally consistent.
161
162func nx_model_spec_validate(s: *NxModelSpec) -> nx_int {
163 if s.n_layers <= 0 { return NX_MS_ERR_BAD_DIM }
164 if s.hidden_dim <= 0 { return NX_MS_ERR_BAD_DIM }
165 if s.n_heads <= 0 { return NX_MS_ERR_BAD_DIM }
166 if s.head_dim <= 0 { return NX_MS_ERR_BAD_DIM }
167 if s.n_kv_heads <= 0 { return NX_MS_ERR_BAD_DIM }
168 if s.ffn_dim <= 0 { return NX_MS_ERR_BAD_DIM }
169 if s.vocab_size <= 0 { return NX_MS_ERR_BAD_VOCAB }
170 if s.max_seq_len <= 0 { return NX_MS_ERR_BAD_DIM }
171 // n_heads * head_dim must equal hidden_dim.
172 if s.n_heads * s.head_dim != s.hidden_dim { return NX_MS_ERR_BAD_HEAD_GEOM }
173 // n_kv_heads must divide n_heads cleanly (GQA constraint).
174 if s.n_heads - (s.n_heads / s.n_kv_heads) * s.n_kv_heads != 0 {
175 return NX_MS_ERR_BAD_HEAD_GEOM
176 }
177 // RoPE needs even head_dim.
178 if s.pos_kind == NX_POS_ROPE {
179 if s.head_dim - (s.head_dim / 2) * 2 != 0 { return NX_MS_ERR_BAD_HEAD_GEOM }
180 }
181 // Sealed enum range gates.
182 if nx_norm_is_valid(s.norm_kind) == 0 { return NX_MS_ERR_BAD_KIND }
183 if nx_act_is_valid(s.act_kind) == 0 { return NX_MS_ERR_BAD_KIND }
184 if nx_pos_is_valid(s.pos_kind) == 0 { return NX_MS_ERR_BAD_KIND }
185 if nx_ffn_is_valid(s.ffn_kind) == 0 { return NX_MS_ERR_BAD_KIND }
186 return NX_MS_OK
187}
188
189// ===== Derived parameter counts ==================================
190//
191// Per-layer parameter count (excluding norms, biases, embeddings).
192// Attention block: 4 matrices of size hidden x (n_heads * head_dim
193// or n_kv_heads * head_dim). FFN block: 3 matrices for SwiGLU,
194// 2 for plain.
195
196func nx_model_spec_attn_params_per_layer(s: *NxModelSpec) -> i64 {
197 let kv_dim: i64 = s.n_kv_heads * s.head_dim
198 // W_q: hidden x hidden; W_k: hidden x kv_dim;
199 // W_v: hidden x kv_dim; W_o: hidden x hidden.
200 return 2 * s.hidden_dim * s.hidden_dim + 2 * s.hidden_dim * kv_dim
201}
202
203func nx_model_spec_ffn_params_per_layer(s: *NxModelSpec) -> i64 {
204 if s.ffn_kind == NX_FFN_SWIGLU { return 3 * s.hidden_dim * s.ffn_dim }
205 if s.ffn_kind == NX_FFN_GEGLU { return 3 * s.hidden_dim * s.ffn_dim }
206 return 2 * s.hidden_dim * s.ffn_dim // PLAIN
207}
208
209func nx_model_spec_total_params(s: *NxModelSpec) -> i64 {
210 let per_layer_attn: i64 = nx_model_spec_attn_params_per_layer(s)
211 let per_layer_ffn: i64 = nx_model_spec_ffn_params_per_layer(s)
212 let per_layer: i64 = per_layer_attn + per_layer_ffn
213 let layers: i64 = s.n_layers * per_layer
214 // Token embedding + output projection.
215 let embed: i64 = s.vocab_size * s.hidden_dim
216 let output: i64 = s.vocab_size * s.hidden_dim
217 return layers + embed + output
218}
219
220// ===== Convenience: known-model factories ========================
221//
222// Reference values per the canonical model cards. Callers can use
223// these as starting points and override.
224
225func nx_model_spec_llama_7b() -> *NxModelSpec {
226 let s: *NxModelSpec = nx_model_spec_new()
227 s.n_layers = 32
228 s.hidden_dim = 4096
229 s.n_heads = 32
230 s.head_dim = 128
231 s.n_kv_heads = 32 // Llama-7B: no GQA
232 s.ffn_dim = 11008
233 s.vocab_size = 32000
234 s.max_seq_len = 2048 // Llama-1 default; Llama-2 extends to 4096
235 s.norm_kind = NX_NORM_RMS
236 s.act_kind = NX_ACT_SILU
237 s.pos_kind = NX_POS_ROPE
238 s.ffn_kind = NX_FFN_SWIGLU
239 return s
240}
241
242func nx_model_spec_mistral_7b() -> *NxModelSpec {
243 let s: *NxModelSpec = nx_model_spec_new()
244 s.n_layers = 32
245 s.hidden_dim = 4096
246 s.n_heads = 32
247 s.head_dim = 128
248 s.n_kv_heads = 8 // Mistral: GQA 4:1
249 s.ffn_dim = 14336
250 s.vocab_size = 32000
251 s.max_seq_len = 32768 // Mistral 7B v0.1 sliding-window
252 s.norm_kind = NX_NORM_RMS
253 s.act_kind = NX_ACT_SILU
254 s.pos_kind = NX_POS_ROPE
255 s.ffn_kind = NX_FFN_SWIGLU
256 return s
257}
258
259func nx_model_spec_gpt2_small() -> *NxModelSpec {
260 let s: *NxModelSpec = nx_model_spec_new()
261 s.n_layers = 12
262 s.hidden_dim = 768
263 s.n_heads = 12
264 s.head_dim = 64
265 s.n_kv_heads = 12
266 s.ffn_dim = 3072
267 s.vocab_size = 50257
268 s.max_seq_len = 1024
269 s.norm_kind = NX_NORM_LAYER // GPT-2 era
270 s.act_kind = NX_ACT_GELU
271 s.pos_kind = NX_POS_LEARNED // absolute learned table
272 s.ffn_kind = NX_FFN_PLAIN // no gating
273 return s
274}
275
276// ===== Self-test ==================================================
277//
278// Closed-form invariants:
279//
280// (a) Default new(): NX_MS_OK
281// (b) Llama-7B factory: NX_MS_OK + total_params ~= 6.74e9
282// (c) Mistral-7B: NX_MS_OK + n_kv_heads=8 (GQA), total params ~= 7.2e9
283// (d) GPT-2 small: NX_MS_OK + LayerNorm + GELU + learned position
284// (e) Validation catches:
285// - n_heads * head_dim != hidden_dim
286// - odd head_dim with RoPE
287// - bad sealed-enum value
288
289func main() -> i64 {
290 // --- (a) Default ---
291 let s: *NxModelSpec = nx_model_spec_new()
292 if nx_model_spec_validate(s) != NX_MS_OK { return 10 }
293
294 // --- (b) Llama-7B ---
295 let llama: *NxModelSpec = nx_model_spec_llama_7b()
296 if nx_model_spec_validate(llama) != NX_MS_OK { return 20 }
297 let llama_params: i64 = nx_model_spec_total_params(llama)
298 // Expected ~6.74B; allow +/- 0.5B for rounding.
299 if llama_params < 6000000000 { return 21 }
300 if llama_params > 7500000000 { return 22 }
301
302 // --- (c) Mistral-7B GQA ---
303 let mistral: *NxModelSpec = nx_model_spec_mistral_7b()
304 if nx_model_spec_validate(mistral) != NX_MS_OK { return 30 }
305 if mistral.n_kv_heads != 8 { return 31 }
306
307 // --- (d) GPT-2 small ---
308 let gpt2: *NxModelSpec = nx_model_spec_gpt2_small()
309 if nx_model_spec_validate(gpt2) != NX_MS_OK { return 40 }
310 if gpt2.norm_kind != NX_NORM_LAYER { return 41 }
311 if gpt2.act_kind != NX_ACT_GELU { return 42 }
312 if gpt2.pos_kind != NX_POS_LEARNED { return 43 }
313
314 // --- (e) Validation rejects bad geometry ---
315 let bad_geom: *NxModelSpec = nx_model_spec_new()
316 bad_geom.hidden_dim = 100 // 32 * 128 = 4096, not 100
317 if nx_model_spec_validate(bad_geom) != NX_MS_ERR_BAD_HEAD_GEOM { return 50 }
318
319 let odd_head: *NxModelSpec = nx_model_spec_new()
320 odd_head.head_dim = 129
321 odd_head.hidden_dim = 129 * 32 // make hidden_dim consistent
322 // pos_kind defaults to ROPE; ROPE needs even head_dim.
323 if nx_model_spec_validate(odd_head) != NX_MS_ERR_BAD_HEAD_GEOM { return 60 }
324
325 let bad_kind: *NxModelSpec = nx_model_spec_new()
326 bad_kind.norm_kind = 99
327 if nx_model_spec_validate(bad_kind) != NX_MS_ERR_BAD_KIND { return 70 }
328
329 // --- (f) Verdict gate ---
330 var vi: nx_int = 0
331 while vi < NX_MS_N_VERDICTS {
332 if nx_ms_verdict_is_valid(vi) != 1 { return 80 + vi }
333 vi = vi + 1
334 }
335
336 return 0
337}