code wiki / (root) / nx_f32_llama_block_v4_test.nx

nx_f32_llama_block_v4_test.nx source

↩ module page · 101 lines · 3848 B

1// nx_f32_llama_block_v4_test.nx -- smoke for nx_f32_llama_block_v4.nx. 2// 3// All-zero F32 lazy weights -> residual identity (out == x). 4// Same identity property as v3 block, just through the dispatcher. 5 6import "nx_syscalls.nx" 7import "nx_tier.nx" 8import "nx_f32.nx" 9import "nx_f32_rmsnorm.nx" 10import "nx_f32_matmul.nx" 11import "nx_f32_activations.nx" 12import "nx_f32_rope.nx" 13import "nx_f32_attn_multi.nx" 14import "nx_f32_kv_cache.nx" 15import "nx_f32_attn_cached.nx" 16import "nx_f32_lazy_weight.nx" 17import "nx_f32_llama_block.nx" 18import "nx_f32_llama_block_v4.nx" 19 20func main() -> i64 { 21 var vi: nx_int = 0 22 while vi < NX_BLK4_N_VERDICTS { 23 if nx_blk4_verdict_is_valid(vi) != 1 { return 5 + vi } 24 vi = vi + 1 25 } 26 27 let hidden_dim: nx_int = 4 28 let n_heads: nx_int = 2 29 let n_kv_heads: nx_int = 2 30 let head_dim: nx_int = 2 31 let ffn_dim: nx_int = 8 32 let n_tokens: nx_int = 1 33 let q_dim: nx_int = n_heads * head_dim 34 let kv_dim: nx_int = n_kv_heads * head_dim 35 36 let layer: *NxF32LlamaLayerLazy = nx_f32_llama_layer_lazy_alloc() 37 layer.gamma_attn = sys_mmap(hidden_dim * 8) as *i64 38 layer.gamma_ffn = sys_mmap(hidden_dim * 8) as *i64 39 40 // Zero-initialized f32 storage for each weight matrix. 41 let s_q: *i64 = sys_mmap(hidden_dim * q_dim * 8) as *i64 42 let s_k: *i64 = sys_mmap(hidden_dim * kv_dim * 8) as *i64 43 let s_v: *i64 = sys_mmap(hidden_dim * kv_dim * 8) as *i64 44 let s_o: *i64 = sys_mmap(q_dim * hidden_dim * 8) as *i64 45 let s_g: *i64 = sys_mmap(hidden_dim * ffn_dim * 8) as *i64 46 let s_u: *i64 = sys_mmap(hidden_dim * ffn_dim * 8) as *i64 47 let s_d: *i64 = sys_mmap(ffn_dim * hidden_dim * 8) as *i64 48 49 layer.W_q = nx_f32_lazy_weight_new_f32(s_q, hidden_dim, q_dim) 50 layer.W_k = nx_f32_lazy_weight_new_f32(s_k, hidden_dim, kv_dim) 51 layer.W_v = nx_f32_lazy_weight_new_f32(s_v, hidden_dim, kv_dim) 52 layer.W_o = nx_f32_lazy_weight_new_f32(s_o, q_dim, hidden_dim) 53 layer.W_gate = nx_f32_lazy_weight_new_f32(s_g, hidden_dim, ffn_dim) 54 layer.W_up = nx_f32_lazy_weight_new_f32(s_u, hidden_dim, ffn_dim) 55 layer.W_down = nx_f32_lazy_weight_new_f32(s_d, ffn_dim, hidden_dim) 56 57 var gi: nx_int = 0 58 while gi < hidden_dim { 59 layer.gamma_attn[gi] = 0x3F800000 60 layer.gamma_ffn[gi] = 0x3F800000 61 gi = gi + 1 62 } 63 64 let cache: *NxF32KVCache = nx_f32_kv_cache_alloc(2, n_kv_heads, 8, head_dim) 65 if cache == (0 as *NxF32KVCache) { return 20 } 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 // 1.0 70 x_in[1] = 0x40000000 // 2.0 71 x_in[2] = 0x40400000 // 3.0 72 x_in[3] = 0x40800000 // 4.0 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_block_forward_v4( 79 x_in, n_tokens, hidden_dim, n_heads, n_kv_heads, head_dim, ffn_dim, 80 layer, cache, 0, eps, attn_scale, rope_base, 0, x_out) 81 if v != NX_BLK4_OK { return 30 + v } 82 83 if x_out[0] != x_in[0] { return 40 } 84 if x_out[1] != x_in[1] { return 41 } 85 if x_out[2] != x_in[2] { return 42 } 86 if x_out[3] != x_in[3] { return 43 } 87 88 nx_f32_kv_cache_advance(cache, n_tokens) 89 if nx_f32_kv_cache_get_seq_len(cache) != 1 { return 50 } 90 91 // ===== Decode forward, apply_rope=1 ===== 92 let x_out2: *i64 = sys_mmap(n_tokens * hidden_dim * 8) as *i64 93 let v2: nx_int = nx_f32_llama_block_forward_v4( 94 x_in, n_tokens, hidden_dim, n_heads, n_kv_heads, head_dim, ffn_dim, 95 layer, cache, 0, eps, attn_scale, rope_base, 1, x_out2) 96 if v2 != NX_BLK4_OK { return 60 + v2 } 97 if x_out2[0] != x_in[0] { return 70 } 98 if x_out2[3] != x_in[3] { return 71 } 99 100 return 0 101}