code wiki / (root) / nx_model_spec.nx

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}