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}