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}