code wiki / (root) / nx_f32_llm_test.nx

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}