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}