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}