nx_f32_llm_test.nx source
↩ module page · 152 lines · 5534 B
1// nx_f32_llm_test.nx -- smoke for nx_f32_llm.nx.
2//
3// Build a 2-layer model with one-hot embed + zero block weights +
4// zero lm_head. Verify:
5// (a) Different token_ids run without segfault
6// (b) logits is all zero (zero lm_head)
7// (c) Argmax returns 0 (first token id)
8// (d) Token-id out-of-range verdict
9
10import "nx_syscalls.nx"
11import "nx_tier.nx"
12import "nx_f32.nx"
13import "nx_f32_kv_cache.nx"
14import "nx_f32_llama_block.nx"
15import "nx_f32_llama_stack.nx"
16import "nx_f32_llm.nx"
17
18func make_zero_layer(hidden_dim: nx_int, q_dim: nx_int,
19 kv_dim: nx_int, ffn_dim: nx_int) -> *NxF32LlamaLayer {
20 let layer: *NxF32LlamaLayer = nx_f32_llama_layer_alloc()
21 layer.gamma_attn = sys_mmap(hidden_dim * 8) as *i64
22 layer.gamma_ffn = sys_mmap(hidden_dim * 8) as *i64
23 layer.W_q = sys_mmap(hidden_dim * q_dim * 8) as *i64
24 layer.W_k = sys_mmap(hidden_dim * kv_dim * 8) as *i64
25 layer.W_v = sys_mmap(hidden_dim * kv_dim * 8) as *i64
26 layer.W_o = sys_mmap(q_dim * hidden_dim * 8) as *i64
27 layer.W_gate = sys_mmap(hidden_dim * ffn_dim * 8) as *i64
28 layer.W_up = sys_mmap(hidden_dim * ffn_dim * 8) as *i64
29 layer.W_down = sys_mmap(ffn_dim * hidden_dim * 8) as *i64
30 var gi: nx_int = 0
31 while gi < hidden_dim {
32 layer.gamma_attn[gi] = 0x3F800000
33 layer.gamma_ffn[gi] = 0x3F800000
34 gi = gi + 1
35 }
36 return layer
37}
38
39func main() -> i64 {
40 var vi: nx_int = 0
41 while vi < NX_F32_LLM_N_VERDICTS {
42 if nx_f32_llm_verdict_is_valid(vi) != 1 { return 5 + vi }
43 vi = vi + 1
44 }
45
46 // Sanity check nx_f32_lt: 1.0 < 2.0 = 1, 2.0 < 1.0 = 0, -1.0 < 1.0 = 1,
47 // 1.0 < -1.0 = 0, -2.0 < -1.0 = 1
48 if nx_f32_lt(0x3F800000, 0x40000000) != 1 { return 10 } // 1 < 2
49 if nx_f32_lt(0x40000000, 0x3F800000) != 0 { return 11 } // 2 < 1
50 if nx_f32_lt(0xBF800000, 0x3F800000) != 1 { return 12 } // -1 < 1
51 if nx_f32_lt(0x3F800000, 0xBF800000) != 0 { return 13 } // 1 < -1
52 if nx_f32_lt(0xC0000000, 0xBF800000) != 1 { return 14 } // -2 < -1
53 if nx_f32_lt(0x3F800000, 0x3F800000) != 0 { return 15 } // 1 < 1
54
55 let hidden_dim: nx_int = 4
56 let n_heads: nx_int = 2
57 let n_kv_heads: nx_int = 2
58 let head_dim: nx_int = 2
59 let ffn_dim: nx_int = 8
60 let vocab_size: nx_int = 6
61 let n_layers: nx_int = 2
62 let q_dim: nx_int = n_heads * head_dim
63 let kv_dim: nx_int = n_kv_heads * head_dim
64
65 let model: *NxF32LlamaModel = nx_f32_llama_model_alloc()
66 model.n_layers = n_layers
67 model.hidden_dim = hidden_dim
68 model.n_heads = n_heads
69 model.n_kv_heads = n_kv_heads
70 model.head_dim = head_dim
71 model.ffn_dim = ffn_dim
72 model.vocab_size = vocab_size
73
74 // Embed: row i has 1.0 at position (i % hidden_dim), 0 elsewhere.
75 model.embed_weights = sys_mmap(vocab_size * hidden_dim * 8) as *i64
76 var vi2: nx_int = 0
77 while vi2 < vocab_size {
78 model.embed_weights[vi2 * hidden_dim + (vi2 - (vi2 / hidden_dim) * hidden_dim)] = 0x3F800000
79 vi2 = vi2 + 1
80 }
81
82 // Build 2 zero layers.
83 model.layers = sys_mmap(n_layers * 8) as *i64
84 model.layers[0] = make_zero_layer(hidden_dim, q_dim, kv_dim, ffn_dim) as i64
85 model.layers[1] = make_zero_layer(hidden_dim, q_dim, kv_dim, ffn_dim) as i64
86
87 // Final gamma_out = 1.
88 model.gamma_out = sys_mmap(hidden_dim * 8) as *i64
89 var gj: nx_int = 0
90 while gj < hidden_dim {
91 model.gamma_out[gj] = 0x3F800000
92 gj = gj + 1
93 }
94
95 // LM head = zero -> logits should be zero.
96 model.lm_head = sys_mmap(hidden_dim * vocab_size * 8) as *i64
97
98 let cache: *NxF32KVCache = nx_f32_kv_cache_alloc(n_layers, n_kv_heads, 16, head_dim)
99
100 let token_ids: *i64 = sys_mmap(3 * 8) as *i64
101 token_ids[0] = 0
102 token_ids[1] = 2
103 token_ids[2] = 5
104
105 let logits: *i64 = sys_mmap(3 * vocab_size * 8) as *i64
106
107 let eps: i64 = 0x322BCC77
108 let attn_scale: i64 = 0x3F3504F3
109 let rope_base: i64 = 0x4548F000
110
111 // Prefill 3 tokens.
112 let v: nx_int = nx_f32_llm_forward(model, token_ids, 3, cache,
113 eps, attn_scale, rope_base, 0, logits)
114 if v != NX_F32_LLM_OK { return 20 + v }
115
116 // With lm_head=0, all logits are zero.
117 var li: nx_int = 0
118 while li < 3 * vocab_size {
119 if logits[li] != 0 { return 30 + li }
120 li = li + 1
121 }
122
123 // Argmax of last token: all zeros -> returns 0 (first index).
124 let amax: nx_int = nx_f32_llm_argmax_last(logits, 3, vocab_size)
125 if amax != 0 { return 60 }
126
127 // Cache advanced by 3.
128 if nx_f32_kv_cache_get_seq_len(cache) != 3 { return 70 }
129
130 // Decode 1 more token.
131 let token1: *i64 = sys_mmap(8) as *i64
132 token1[0] = 1
133 let logits2: *i64 = sys_mmap(vocab_size * 8) as *i64
134 let v2: nx_int = nx_f32_llm_forward(model, token1, 1, cache,
135 eps, attn_scale, rope_base, 0, logits2)
136 if v2 != NX_F32_LLM_OK { return 80 + v2 }
137 var lj: nx_int = 0
138 while lj < vocab_size {
139 if logits2[lj] != 0 { return 90 + lj }
140 lj = lj + 1
141 }
142 if nx_f32_kv_cache_get_seq_len(cache) != 4 { return 100 }
143
144 // Token id out of range -> ERR_TOKEN.
145 let token_bad: *i64 = sys_mmap(8) as *i64
146 token_bad[0] = 99
147 let v3: nx_int = nx_f32_llm_forward(model, token_bad, 1, cache,
148 eps, attn_scale, rope_base, 0, logits2)
149 if v3 != NX_F32_LLM_ERR_TOKEN { return 110 }
150
151 return 0
152}