code wiki / (root) / nx_f32_llama_stack_v4.nx

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}