code wiki / (root) / nx_transformer_stack.nx

nx_transformer_stack.nx source

↩ module page · 172 lines · 6889 B

1// nx_transformer_stack.nx -- multi-layer transformer forward composer. 2// 3// L5 brick. The missing layer between "load weights for one block" 4// (nx_gguf_load_block) and "run the full LLM pipeline" (nx_llm_run 5// v2 -- next brick). 6// 7// Composes: 8// - Pre-allocated scratch tensors (one set, reused across layers) 9// - Per-layer loop that: 10// * loads block weights from GGUF via nx_gguf_load_block_weights 11// * calls nx_transformer_block_forward with the loaded bundle 12// - Optional final RMSNorm (output_norm.weight from load_model) 13// 14// Does NOT include: embed, output_projection, sample. Those compose 15// in nx_llm_run v2. This brick proves multi-layer chain works. 16// 17// Bits-up composition (every primitive already canonical): 18// nx_tensor.nx -- NxTensor + alloc 19// nx_rmsnorm.nx -- final norm 20// nx_transformer_block.nx -- per-layer forward + weight bundle 21// nx_gguf.nx + nx_gguf_load.nx -- header + per-tensor load 22// nx_gguf_load_block.nx -- 9-tensor per-layer loader 23// 24// genealogy_id: standard_transformer_stack_pattern + 25// touvron_2023_llama_decoder 26// lineage_id: substrate_transformer_stack_v1_single_head 27 28// nx_safety_envelope: 29// intended_use: "Run N-layer transformer decoder forward 30// sourcing each layer's weights from a parsed 31// Llama-class GGUF; single-head v1; allocates 32// scratch tensors once and reuses across layers" 33// sil_target: SIL2 34// asil_target: QM 35// dal_target: DAL C 36// evidence: [bounded_layer_loop, composes_only_shipped, 37// scratch_reuse_proven_pattern] 38// hazard_register: [bug-tape-scratch-shape-mismatch-vs-spec, 39// bug-tape-layer-weight-leak-across-iterations, 40// bug-tape-attn-scale-default-mismatch] 41// residual_risk: "Numerical correctness verified by component 42// smokes; integration correctness vs reference 43// is the bench-vs-llama.cpp task" 44// verdict: NOT_YET_EVALUATED 45 46import "nx_syscalls.nx" 47import "nx_tier.nx" 48import "nx_loop.nx" 49import "nx_tensor.nx" 50import "nx_rmsnorm.nx" 51import "nx_gguf.nx" 52import "nx_gguf_load.nx" 53import "nx_gguf_load_block.nx" 54import "nx_transformer_block.nx" 55 56// ===== Sealed-enum: StackVerdict ================================== 57 58const NX_TS_OK: nx_int = 0 59const NX_TS_ERR_BAD_DIM: nx_int = 1 60const NX_TS_ERR_OOM: nx_int = 2 61const NX_TS_ERR_LAYER_LOAD: nx_int = 3 62const NX_TS_ERR_BLOCK_FWD: nx_int = 4 63const NX_TS_ERR_NORM_FAIL: nx_int = 5 64const NX_TS_N_VERDICTS: nx_int = 6 65 66func nx_ts_verdict_is_valid(v: nx_int) -> nx_int { 67 if v < 0 { return 0 } 68 if v >= NX_TS_N_VERDICTS { return 0 } 69 return 1 70} 71 72// Helper: allocate an [n_tokens, dim] Q10 tensor. 73func _ts_alloc_2d(n_tokens: nx_int, dim: nx_int) -> *NxTensor { 74 let sh: *i64 = sys_mmap(2 * 8) as *i64 75 sh[0] = n_tokens 76 sh[1] = dim 77 let err: *i64 = sys_mmap(8) as *i64 78 err[0] = 0 79 let t: *NxTensor = nx_t_alloc(NX_DT_I64, sh, 2, err) 80 return t 81} 82 83// ===== Public: multi-layer transformer forward ==================== 84// 85// x: [n_tokens, hidden_dim] Q10 input embedded sequence 86// positions: [n_tokens] token positions for RoPE 87// buf: GGUF byte buffer 88// hdr: parsed GGUF header 89// n_layers: number of transformer blocks to run 90// hidden_dim: embedding width (must match GGUF + x.shape[1]) 91// head_dim: per-head dimension (single-head v1: == hidden_dim) 92// ffn_dim: FFN intermediate dimension 93// rope_base: RoPE theta base (10000 / 500000) 94// attn_scale_q10: sqrt(head_dim) inverse, in Q10 (caller computes) 95// final_norm_gamma: *i64 [hidden_dim] -- if non-null, apply RMSNorm 96// to x after all layers 97// 98// 11 args -- under the 16-arg limit. 99 100func nx_transformer_stack_forward( 101 x: *NxTensor, 102 positions: *i64, 103 buf: *u8, hdr: *NxGgufHeader, 104 n_layers: nx_int, 105 hidden_dim: nx_int, 106 head_dim: nx_int, 107 ffn_dim: nx_int, 108 rope_base: nx_int, 109 attn_scale_q10: nx_int, 110 final_norm_gamma: *i64) -> nx_int { 111 112 if n_layers <= 0 { return NX_TS_ERR_BAD_DIM } 113 if hidden_dim <= 0 { return NX_TS_ERR_BAD_DIM } 114 if head_dim <= 0 { return NX_TS_ERR_BAD_DIM } 115 if ffn_dim <= 0 { return NX_TS_ERR_BAD_DIM } 116 if x.dtype != NX_DT_I64 { return NX_TS_ERR_BAD_DIM } 117 if x.ndim != 2 { return NX_TS_ERR_BAD_DIM } 118 if x.shape[1] != hidden_dim { return NX_TS_ERR_BAD_DIM } 119 120 let n_tok: nx_int = x.shape[0] 121 122 // ----- Allocate scratch tensors (once; reused across layers) ----- 123 let s_norm: *NxTensor = _ts_alloc_2d(n_tok, hidden_dim) 124 let s_q: *NxTensor = _ts_alloc_2d(n_tok, head_dim) 125 let s_k: *NxTensor = _ts_alloc_2d(n_tok, head_dim) 126 let s_v: *NxTensor = _ts_alloc_2d(n_tok, head_dim) 127 let s_attn_out: *NxTensor = _ts_alloc_2d(n_tok, hidden_dim) 128 let s_attn_proj: *NxTensor = _ts_alloc_2d(n_tok, hidden_dim) 129 let s_ffn_gate: *NxTensor = _ts_alloc_2d(n_tok, ffn_dim) 130 let s_ffn_up: *NxTensor = _ts_alloc_2d(n_tok, ffn_dim) 131 let s_ffn_hidden: *NxTensor = _ts_alloc_2d(n_tok, ffn_dim) 132 let s_ffn_proj: *NxTensor = _ts_alloc_2d(n_tok, hidden_dim) 133 134 if s_norm == (0 as *NxTensor) { return NX_TS_ERR_OOM } 135 if s_ffn_proj == (0 as *NxTensor) { return NX_TS_ERR_OOM } 136 137 // ----- Allocate one block weight bundle (reused) ----- 138 let w: *NxTransformerBlockWeights = 139 sys_mmap(NX_TB_WEIGHTS_BYTES) as *NxTransformerBlockWeights 140 let err: *i64 = sys_mmap(8) as *i64 141 142 // ----- Per-layer loop ----- 143 var layer: nx_int = 0 144 var iter: nx_int = 0 145 var verdict: nx_int = NX_LOOP_RUNNING 146 let BUDGET: nx_int = n_layers 147 while verdict == NX_LOOP_RUNNING && iter < BUDGET { 148 err[0] = 0 149 let v_load: nx_int = nx_gguf_load_block_weights( 150 buf, hdr, layer, head_dim, rope_base, w, err) 151 if v_load != NX_GBL_OK { return NX_TS_ERR_LAYER_LOAD } 152 153 let v_fwd: nx_int = nx_transformer_block_forward( 154 x, positions, w, 155 s_norm, s_q, s_k, s_v, 156 s_attn_out, s_attn_proj, 157 s_ffn_gate, s_ffn_up, s_ffn_hidden, s_ffn_proj, 158 attn_scale_q10) 159 if v_fwd != NX_TB_OK { return NX_TS_ERR_BLOCK_FWD } 160 161 layer = layer + 1 162 iter = iter + 1 163 } 164 165 // ----- Final RMSNorm (optional) ----- 166 if final_norm_gamma != (0 as *i64) { 167 let v_n: nx_int = nx_rmsnorm_forward(x, final_norm_gamma, x) 168 if v_n != NX_RMSN_OK { return NX_TS_ERR_NORM_FAIL } 169 } 170 171 return NX_TS_OK 172}