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}