code wiki / (root) / nx_f32_llm_v4_test.nx

nx_f32_llm_v4_test.nx source

↩ module page · 89 lines · 3101 B

1// nx_f32_llm_v4_test.nx -- smoke for the lazy LLM forward. 2 3import "nx_syscalls.nx" 4import "nx_tier.nx" 5import "nx_gguf.nx" 6import "nx_gguf_fixture_tiny.nx" 7import "nx_gguf_load_f32.nx" 8import "nx_f32_kv_cache.nx" 9import "nx_f32_lazy_weight.nx" 10import "nx_f32_llama_block_v4.nx" 11import "nx_f32_llama_stack_v4.nx" 12import "nx_f32_llama_layer_lazy_load.nx" 13import "nx_f32_llm.nx" 14import "nx_f32_llm_v4.nx" 15 16func main() -> i64 { 17 var vi: nx_int = 0 18 while vi < NX_FLV4_N_VERDICTS { 19 if nx_flv4_verdict_is_valid(vi) != 1 { return 5 + vi } 20 vi = vi + 1 21 } 22 23 let b: *NxGgufFixtureBundle = nx_gft_build_tiny_llama(42 as i64) 24 if nx_gft_is_built(b) != 1 { return 10 } 25 26 let model: *NxF32LlamaModel = nx_f32_llama_model_alloc() 27 model.n_layers = b.n_layers 28 model.hidden_dim = b.hidden_dim 29 model.n_heads = b.n_heads 30 model.n_kv_heads = b.n_heads 31 model.head_dim = b.head_dim 32 model.ffn_dim = b.ffn_dim 33 model.vocab_size = b.vocab_size 34 35 let out_err: *i64 = sys_mmap(8) as *i64 36 let v_load: nx_int = nx_f32_llm_load_weights_v4_from_gguf(b.gguf_buf, b.hdr, 37 model, out_err) 38 if v_load != NX_FLV4_OK { return 20 + v_load } 39 40 // All four top-level non-null. 41 if (model.embed_weights as i64) == 0 { return 40 } 42 if (model.gamma_out as i64) == 0 { return 41 } 43 if (model.lm_head as i64) == 0 { return 42 } 44 if (model.layers as i64) == 0 { return 43 } 45 46 // Layer 0 is a NxF32LlamaLayerLazy. 47 let layer0: *NxF32LlamaLayerLazy = (model.layers[0]) as *NxF32LlamaLayerLazy 48 if (layer0 as i64) == 0 { return 50 } 49 if (layer0.W_q as i64) == 0 { return 51 } 50 if layer0.W_q.dtype_tag != NX_LW_DTYPE_F32 { return 52 } 51 52 // Forward pass. 53 let cache: *NxF32KVCache = nx_f32_kv_cache_alloc( 54 model.n_layers, model.n_kv_heads, 8, model.head_dim) 55 if cache == (0 as *NxF32KVCache) { return 60 } 56 57 let token_ids: *i64 = sys_mmap(2 * 8) as *i64 58 token_ids[0] = 0 59 token_ids[1] = 1 60 61 let logits: *i64 = sys_mmap(2 * model.vocab_size * 8) as *i64 62 63 let eps: i64 = 0x322BCC77 64 let attn_scale: i64 = 0x3F3504F3 65 let rope_base: i64 = 0x4548F000 66 67 let v: nx_int = nx_f32_llm_forward_v4(model, token_ids, 2, cache, 68 eps, attn_scale, rope_base, 0, logits) 69 if v != NX_FLV4_OK { return 70 + v } 70 71 // Zero-weight fixture -> all logits zero. 72 var li: nx_int = 0 73 while li < 2 * model.vocab_size { 74 if logits[li] != 0 { return 90 + li } 75 li = li + 1 76 } 77 if nx_f32_kv_cache_get_seq_len(cache) != 2 { return 100 } 78 79 // Decode 1 more token. 80 let next_id: *i64 = sys_mmap(8) as *i64 81 next_id[0] = 0 82 let logits2: *i64 = sys_mmap(model.vocab_size * 8) as *i64 83 let v2: nx_int = nx_f32_llm_forward_v4(model, next_id, 1, cache, 84 eps, attn_scale, rope_base, 0, logits2) 85 if v2 != NX_FLV4_OK { return 110 + v2 } 86 if nx_f32_kv_cache_get_seq_len(cache) != 3 { return 120 } 87 88 return 0 89}