code wiki / (root) / nx_f32_llama_stack_v4_test.nx

nx_f32_llama_stack_v4_test.nx source

↩ module page · 103 lines · 3958 B

1// nx_f32_llama_stack_v4_test.nx -- smoke for v4 stack. 2 3import "nx_syscalls.nx" 4import "nx_tier.nx" 5import "nx_f32.nx" 6import "nx_f32_kv_cache.nx" 7import "nx_f32_lazy_weight.nx" 8import "nx_f32_llama_block_v4.nx" 9import "nx_f32_llama_stack_v4.nx" 10 11func _make_zero_lazy_layer(hidden_dim: nx_int, q_dim: nx_int, 12 kv_dim: nx_int, ffn_dim: nx_int) -> *NxF32LlamaLayerLazy { 13 let layer: *NxF32LlamaLayerLazy = nx_f32_llama_layer_lazy_alloc() 14 layer.gamma_attn = sys_mmap(hidden_dim * 8) as *i64 15 layer.gamma_ffn = sys_mmap(hidden_dim * 8) as *i64 16 17 let s_q: *i64 = sys_mmap(hidden_dim * q_dim * 8) as *i64 18 let s_k: *i64 = sys_mmap(hidden_dim * kv_dim * 8) as *i64 19 let s_v: *i64 = sys_mmap(hidden_dim * kv_dim * 8) as *i64 20 let s_o: *i64 = sys_mmap(q_dim * hidden_dim * 8) as *i64 21 let s_g: *i64 = sys_mmap(hidden_dim * ffn_dim * 8) as *i64 22 let s_u: *i64 = sys_mmap(hidden_dim * ffn_dim * 8) as *i64 23 let s_d: *i64 = sys_mmap(ffn_dim * hidden_dim * 8) as *i64 24 25 layer.W_q = nx_f32_lazy_weight_new_f32(s_q, hidden_dim, q_dim) 26 layer.W_k = nx_f32_lazy_weight_new_f32(s_k, hidden_dim, kv_dim) 27 layer.W_v = nx_f32_lazy_weight_new_f32(s_v, hidden_dim, kv_dim) 28 layer.W_o = nx_f32_lazy_weight_new_f32(s_o, q_dim, hidden_dim) 29 layer.W_gate = nx_f32_lazy_weight_new_f32(s_g, hidden_dim, ffn_dim) 30 layer.W_up = nx_f32_lazy_weight_new_f32(s_u, hidden_dim, ffn_dim) 31 layer.W_down = nx_f32_lazy_weight_new_f32(s_d, ffn_dim, hidden_dim) 32 33 var gi: nx_int = 0 34 while gi < hidden_dim { 35 layer.gamma_attn[gi] = 0x3F800000 36 layer.gamma_ffn[gi] = 0x3F800000 37 gi = gi + 1 38 } 39 return layer 40} 41 42func main() -> i64 { 43 var vi: nx_int = 0 44 while vi < NX_STK4_N_VERDICTS { 45 if nx_stk4_verdict_is_valid(vi) != 1 { return 5 + vi } 46 vi = vi + 1 47 } 48 49 let hidden_dim: nx_int = 4 50 let n_heads: nx_int = 2 51 let n_kv_heads: nx_int = 2 52 let head_dim: nx_int = 2 53 let ffn_dim: nx_int = 8 54 let n_tokens: nx_int = 1 55 let n_layers: nx_int = 3 56 let q_dim: nx_int = n_heads * head_dim 57 let kv_dim: nx_int = n_kv_heads * head_dim 58 59 let layers: *i64 = sys_mmap(n_layers * 8) as *i64 60 layers[0] = _make_zero_lazy_layer(hidden_dim, q_dim, kv_dim, ffn_dim) as i64 61 layers[1] = _make_zero_lazy_layer(hidden_dim, q_dim, kv_dim, ffn_dim) as i64 62 layers[2] = _make_zero_lazy_layer(hidden_dim, q_dim, kv_dim, ffn_dim) as i64 63 64 let cache: *NxF32KVCache = nx_f32_kv_cache_alloc(n_layers, n_kv_heads, 8, head_dim) 65 if cache == (0 as *NxF32KVCache) { return 10 } 66 67 let x_in: *i64 = sys_mmap(n_tokens * hidden_dim * 8) as *i64 68 let x_out: *i64 = sys_mmap(n_tokens * hidden_dim * 8) as *i64 69 x_in[0] = 0x3F800000 70 x_in[1] = 0x40000000 71 x_in[2] = 0x40400000 72 x_in[3] = 0x40800000 73 74 let eps: i64 = 0x322BCC77 75 let attn_scale: i64 = 0x3F3504F3 76 let rope_base: i64 = 0x4548F000 77 78 let v: nx_int = nx_f32_llama_stack_forward_v4( 79 x_in, x_out, n_tokens, n_layers, 80 hidden_dim, n_heads, n_kv_heads, head_dim, ffn_dim, 81 layers, cache, eps, attn_scale, rope_base, 0) 82 if v != NX_STK4_OK { return 20 + v } 83 84 if x_out[0] != x_in[0] { return 30 } 85 if x_out[1] != x_in[1] { return 31 } 86 if x_out[2] != x_in[2] { return 32 } 87 if x_out[3] != x_in[3] { return 33 } 88 89 if nx_f32_kv_cache_get_seq_len(cache) != 1 { return 40 } 90 91 // Second forward (decode style). 92 let x_out2: *i64 = sys_mmap(n_tokens * hidden_dim * 8) as *i64 93 let v2: nx_int = nx_f32_llama_stack_forward_v4( 94 x_in, x_out2, n_tokens, n_layers, 95 hidden_dim, n_heads, n_kv_heads, head_dim, ffn_dim, 96 layers, cache, eps, attn_scale, rope_base, 1) 97 if v2 != NX_STK4_OK { return 50 + v2 } 98 if x_out2[0] != x_in[0] { return 60 } 99 if x_out2[3] != x_in[3] { return 61 } 100 if nx_f32_kv_cache_get_seq_len(cache) != 2 { return 70 } 101 102 return 0 103}