code wiki / (root) / nx_f32_llama_stack_test.nx

nx_f32_llama_stack_test.nx source

↩ module page · 120 lines · 4744 B

1// nx_f32_llama_stack_test.nx -- smoke for nx_f32_llama_stack.nx. 2// 3// Build a 3-layer stack with all-zero weights and gamma=1. 4// Forward through all 3 layers should preserve x_in bit-exact 5// (residual passes through every layer, all projections are zero). 6 7import "nx_syscalls.nx" 8import "nx_tier.nx" 9import "nx_f32.nx" 10import "nx_f32_kv_cache.nx" 11import "nx_f32_llama_block.nx" 12import "nx_f32_llama_stack.nx" 13 14func make_zero_layer(hidden_dim: nx_int, q_dim: nx_int, 15 kv_dim: nx_int, ffn_dim: nx_int) -> *NxF32LlamaLayer { 16 let layer: *NxF32LlamaLayer = nx_f32_llama_layer_alloc() 17 layer.gamma_attn = sys_mmap(hidden_dim * 8) as *i64 18 layer.gamma_ffn = sys_mmap(hidden_dim * 8) as *i64 19 layer.W_q = sys_mmap(hidden_dim * q_dim * 8) as *i64 20 layer.W_k = sys_mmap(hidden_dim * kv_dim * 8) as *i64 21 layer.W_v = sys_mmap(hidden_dim * kv_dim * 8) as *i64 22 layer.W_o = sys_mmap(q_dim * hidden_dim * 8) as *i64 23 layer.W_gate = sys_mmap(hidden_dim * ffn_dim * 8) as *i64 24 layer.W_up = sys_mmap(hidden_dim * ffn_dim * 8) as *i64 25 layer.W_down = sys_mmap(ffn_dim * hidden_dim * 8) as *i64 26 27 var gi: nx_int = 0 28 while gi < hidden_dim { 29 layer.gamma_attn[gi] = 0x3F800000 // 1.0 30 layer.gamma_ffn[gi] = 0x3F800000 31 gi = gi + 1 32 } 33 return layer 34} 35 36func main() -> i64 { 37 var vi: nx_int = 0 38 while vi < NX_F32_STK_N_VERDICTS { 39 if nx_f32_stk_verdict_is_valid(vi) != 1 { return 5 + vi } 40 vi = vi + 1 41 } 42 43 let hidden_dim: nx_int = 4 44 let n_heads: nx_int = 2 45 let n_kv_heads: nx_int = 2 46 let head_dim: nx_int = 2 47 let ffn_dim: nx_int = 8 48 let n_tokens: nx_int = 1 49 let n_layers: nx_int = 3 50 let q_dim: nx_int = n_heads * head_dim 51 let kv_dim: nx_int = n_kv_heads * head_dim 52 53 // Build 3 zero-weight layers; pack pointers into i64 array. 54 let layers: *i64 = sys_mmap(n_layers * 8) as *i64 55 layers[0] = make_zero_layer(hidden_dim, q_dim, kv_dim, ffn_dim) as i64 56 layers[1] = make_zero_layer(hidden_dim, q_dim, kv_dim, ffn_dim) as i64 57 layers[2] = make_zero_layer(hidden_dim, q_dim, kv_dim, ffn_dim) as i64 58 59 let cache: *NxF32KVCache = nx_f32_kv_cache_alloc(n_layers, n_kv_heads, 8, head_dim) 60 if cache == (0 as *NxF32KVCache) { return 10 } 61 62 let x_in: *i64 = sys_mmap(n_tokens * hidden_dim * 8) as *i64 63 let x_out: *i64 = sys_mmap(n_tokens * hidden_dim * 8) as *i64 64 x_in[0] = 0x3F800000 // 1.0 65 x_in[1] = 0x40000000 // 2.0 66 x_in[2] = 0x40400000 // 3.0 67 x_in[3] = 0x40800000 // 4.0 68 69 let eps: i64 = 0x322BCC77 70 let attn_scale: i64 = 0x3F3504F3 71 let rope_base: i64 = 0x4548F000 72 73 // First forward (prefill-style, fresh cache). 74 let v: nx_int = nx_f32_llama_stack_forward( 75 x_in, x_out, n_tokens, n_layers, 76 hidden_dim, n_heads, n_kv_heads, head_dim, ffn_dim, 77 layers, cache, eps, attn_scale, rope_base, 0) 78 if v != NX_F32_STK_OK { return 20 + v } 79 80 // All-zero weights -> residual passes through every layer unchanged. 81 if x_out[0] != x_in[0] { return 30 } 82 if x_out[1] != x_in[1] { return 31 } 83 if x_out[2] != x_in[2] { return 32 } 84 if x_out[3] != x_in[3] { return 33 } 85 86 // Cache advanced by n_tokens after the stack call. 87 if nx_f32_kv_cache_get_seq_len(cache) != 1 { return 40 } 88 89 // Second forward (decode-style, cache.seq_len was 1). 90 let x_out2: *i64 = sys_mmap(n_tokens * hidden_dim * 8) as *i64 91 let v2: nx_int = nx_f32_llama_stack_forward( 92 x_in, x_out2, n_tokens, n_layers, 93 hidden_dim, n_heads, n_kv_heads, head_dim, ffn_dim, 94 layers, cache, eps, attn_scale, rope_base, 0) 95 if v2 != NX_F32_STK_OK { return 50 + v2 } 96 if x_out2[0] != x_in[0] { return 60 } 97 if x_out2[3] != x_in[3] { return 61 } 98 99 if nx_f32_kv_cache_get_seq_len(cache) != 2 { return 70 } 100 101 // Test RoPE-enabled path. Q=K=0 from zero weights so RoPE is a no-op. 102 let x_out3: *i64 = sys_mmap(n_tokens * hidden_dim * 8) as *i64 103 let v3: nx_int = nx_f32_llama_stack_forward( 104 x_in, x_out3, n_tokens, n_layers, 105 hidden_dim, n_heads, n_kv_heads, head_dim, ffn_dim, 106 layers, cache, eps, attn_scale, rope_base, 1) 107 if v3 != NX_F32_STK_OK { return 80 + v3 } 108 if x_out3[0] != x_in[0] { return 90 } 109 if x_out3[3] != x_in[3] { return 91 } 110 if nx_f32_kv_cache_get_seq_len(cache) != 3 { return 100 } 111 112 // Bad-dim verdict. 113 let vE: nx_int = nx_f32_llama_stack_forward( 114 x_in, x_out, 0, n_layers, 115 hidden_dim, n_heads, n_kv_heads, head_dim, ffn_dim, 116 layers, cache, eps, attn_scale, rope_base, 0) 117 if vE != NX_F32_STK_ERR_BAD_DIM { return 110 } 118 119 return 0 120}