nx_f32_llama_stack_v4.nx source
↩ module page · 96 lines · 3270 B
1// nx_f32_llama_stack_v4.nx -- N-layer transformer stack using v4 lazy layers.
2//
3// Mirrors nx_f32_llama_stack but accepts an array of *NxF32LlamaLayerLazy
4// and dispatches each layer through nx_f32_llama_block_forward_v4.
5//
6// genealogy_id: standard_transformer_stack + lazy_dispatch
7// lineage_id: substrate_f32_llama_stack_v4
8
9import "nx_syscalls.nx"
10import "nx_tier.nx"
11import "nx_f32.nx"
12import "nx_f32_kv_cache.nx"
13import "nx_f32_lazy_weight.nx"
14import "nx_f32_llama_block_v4.nx"
15
16const NX_STK4_OK: nx_int = 0
17const NX_STK4_ERR_BAD_DIM: nx_int = 1
18const NX_STK4_ERR_NULL: nx_int = 2
19const NX_STK4_ERR_LAYER: nx_int = 3
20const NX_STK4_N_VERDICTS: nx_int = 4
21
22func nx_stk4_verdict_is_valid(v: nx_int) -> nx_int {
23 if v < 0 { return 0 }
24 if v >= NX_STK4_N_VERDICTS { return 0 }
25 return 1
26}
27
28// 15 args -- at the limit.
29
30func nx_f32_llama_stack_forward_v4(
31 x_in: *i64,
32 x_out: *i64,
33 n_tokens: nx_int,
34 n_layers: nx_int,
35 hidden_dim: nx_int,
36 n_heads: nx_int,
37 n_kv_heads: nx_int,
38 head_dim: nx_int,
39 ffn_dim: nx_int,
40 layers: *i64,
41 cache: *NxF32KVCache,
42 eps: i64,
43 attn_scale: i64,
44 rope_log_base: i64,
45 apply_rope: nx_int) -> nx_int {
46
47 if n_tokens <= 0 { return NX_STK4_ERR_BAD_DIM }
48 if n_layers <= 0 { return NX_STK4_ERR_BAD_DIM }
49 if hidden_dim <= 0 { return NX_STK4_ERR_BAD_DIM }
50 if x_in == (0 as *i64) { return NX_STK4_ERR_NULL }
51 if x_out == (0 as *i64) { return NX_STK4_ERR_NULL }
52 if layers == (0 as *i64) { return NX_STK4_ERR_NULL }
53 if cache == (0 as *NxF32KVCache) { return NX_STK4_ERR_NULL }
54
55 let n_elem: nx_int = n_tokens * hidden_dim
56 let buf_a: *i64 = sys_mmap(n_elem * 8) as *i64
57 let buf_b: *i64 = sys_mmap(n_elem * 8) as *i64
58
59 var i: nx_int = 0
60 while i < n_elem { buf_a[i] = x_in[i]; i = i + 1 }
61
62 var L: nx_int = 0
63 var use_a_as_src: nx_int = 1
64 while L < n_layers {
65 let layer_ptr_i64: i64 = layers[L]
66 let layer: *NxF32LlamaLayerLazy = layer_ptr_i64 as *NxF32LlamaLayerLazy
67 if layer == (0 as *NxF32LlamaLayerLazy) { return NX_STK4_ERR_LAYER }
68
69 if use_a_as_src != 0 {
70 let v: nx_int = nx_f32_llama_block_forward_v4(
71 buf_a, n_tokens, hidden_dim, n_heads, n_kv_heads, head_dim, ffn_dim,
72 layer, cache, L, eps, attn_scale, rope_log_base, apply_rope, buf_b)
73 if v != NX_BLK4_OK { return NX_STK4_ERR_LAYER }
74 use_a_as_src = 0
75 } else {
76 let v: nx_int = nx_f32_llama_block_forward_v4(
77 buf_b, n_tokens, hidden_dim, n_heads, n_kv_heads, head_dim, ffn_dim,
78 layer, cache, L, eps, attn_scale, rope_log_base, apply_rope, buf_a)
79 if v != NX_BLK4_OK { return NX_STK4_ERR_LAYER }
80 use_a_as_src = 1
81 }
82 L = L + 1
83 }
84
85 nx_f32_kv_cache_advance(cache, n_tokens)
86
87 if use_a_as_src != 0 {
88 var k: nx_int = 0
89 while k < n_elem { x_out[k] = buf_a[k]; k = k + 1 }
90 } else {
91 var k2: nx_int = 0
92 while k2 < n_elem { x_out[k2] = buf_b[k2]; k2 = k2 + 1 }
93 }
94
95 return NX_STK4_OK
96}