code wiki / (root) / nx_llm_argmax_probe.nx

nx_llm_argmax_probe.nx source

↩ module page · 93 lines · 4481 B

1// nx_llm_argmax_probe.nx -- PURE-FORWARD correctness test. Bypasses run_v3's sampler + repetition-penalty: 2// manual sequential prefill of "The capital of France is", then reads the RAW top-5 logits of the final position 3// and decodes them. If top-1 = " Paris" (or coherent), the FORWARD is correct and the "are are are" loop was a 4// SAMPLER artifact (debt in nx_f32_sampler/run_v3), not the forward. If top-1 is garbage, the forward residual is 5// real. Decisive isolation of forward-vs-sampler. expect_exit: 0 6import "nx_syscalls.nx" 7import "nx_tier.nx" 8import "nx_bpe.nx" 9import "nx_gguf.nx" 10import "nx_gguf_load.nx" 11import "nx_gguf_meta.nx" 12import "nx_f32.nx" 13import "nx_f32_kv_cache.nx" 14import "nx_f32_lazy_weight.nx" 15import "nx_f32_llama_block.nx" 16import "nx_f32_llama_block_v4.nx" 17import "nx_f32_llama_stack_v4.nx" 18import "nx_f32_llama_layer_lazy_load.nx" 19import "nx_f32_llm.nx" 20import "nx_f32_llm_v4.nx" 21import "nx_f32_llm_read_dims.nx" 22import "nx_f32_bpe_load.nx" 23import "nx_f32_llm_special_tokens.nx" 24 25func pr_puts(s: *u8) -> i64 { var n: i64=0; while s[n]!=(0 as u8){n=n+1} sys_write(1,s,n); return 0 } 26func pr_num(v: i64) -> i64 { let bb: *u8=sys_mmap(28); var m: i64=v; if m<0{sys_write(1,"-" as *u8,1);m=0-m} let t: *u8=sys_mmap(28); var k: i64=0; if m==0{t[0]=48 as u8;k=1} while m>0{t[k]=(48+(m%10)) as u8;m=m/10;k=k+1} var i: i64=0; while i<k{bb[i]=t[k-1-i];i=i+1} sys_write(1,bb,k); return 0 } 27 28func main() -> i64 { 29 let path: *u8 = "/tmp/nx_real_model.gguf" as *u8 30 let len_out: *i64 = sys_mmap(8) as *i64 31 let buf: *u8 = sys_read_file(path, len_out) 32 if buf == (0 as *u8) { return 10 } 33 let hdr: *NxGgufHeader = sys_mmap(NX_GGUF_HDR_BYTES) as *NxGgufHeader 34 if nx_gguf_parse(buf, len_out[0], hdr) != NX_GGUF_OK { return 20 } 35 let model: *NxF32LlamaModel = nx_f32_llama_model_alloc() 36 let out_err: *i64 = sys_mmap(8) as *i64 37 if nx_f32_llm_read_dims_from_gguf(buf, len_out[0], hdr, model, out_err) != NX_FLD_OK { return 30 } 38 if nx_f32_llm_load_weights_v4_from_gguf(buf, hdr, model, out_err) != NX_FLV4_OK { return 40 } 39 let vocab: *NxBpeVocab = nx_bpe_vocab_new(67108864, 262144, 524288) 40 let nt: *i64 = sys_mmap(8) as *i64 41 let nm: *i64 = sys_mmap(8) as *i64 42 if nx_f32_bpe_load_from_gguf(buf, len_out[0], hdr, vocab, nt, nm, out_err) != NX_FBL_OK { return 50 } 43 44 let eps: i64 = 0x358637BD 45 let attn_scale: i64 = 0x3E000000 46 let rope_base: i64 = 0x415D0EAB 47 let V: nx_int = model.vocab_size 48 49 // manual sequential prefill, keep last logits 50 let toks: *i64 = sys_mmap(16 * 8) as *i64 51 toks[0]=785; toks[1]=6722; toks[2]=315; toks[3]=9625; toks[4]=374 // "The capital of France is" 52 let np: nx_int = 5 53 let cache: *NxF32KVCache = nx_f32_kv_cache_alloc(model.n_layers, model.n_kv_heads, 64, model.head_dim) 54 let logits: *i64 = sys_mmap(V * 8) as *i64 55 var p: nx_int = 0 56 while p < np { 57 let one: *i64 = (((toks as i64) + p * 8)) as *i64 58 if nx_f32_llm_forward_v4(model, one, 1, cache, eps, attn_scale, rope_base, 1, logits) != NX_FLV4_OK { return 60 } 59 p = p + 1 60 } 61 62 pr_puts("=== PURE-FORWARD argmax for 'The capital of France is' (no sampler) ===\n" as *u8) 63 // top-5 by repeated max-scan (mask picked with -inf) 64 let NEG: i64 = 0xFF800000 65 var r: nx_int = 0 66 while r < 5 { 67 var best: nx_int = 0 68 var bestv: i64 = logits[0] 69 var i: nx_int = 1 70 while i < V { 71 let lv: i64 = logits[i] & 0xFFFFFFFF 72 // f32 compare: bestv vs lv (both may be neg). compare as f32 via sign+magnitude. 73 let a: i64 = bestv & 0xFFFFFFFF 74 var gt: nx_int = 0 75 let sa: i64 = (a >> 31) & 1 76 let sb: i64 = (lv >> 31) & 1 77 if sa == 1 { if sb == 1 { if lv < a { gt = 1 } } else { gt = 1 } } 78 else { if sb == 0 { if lv > a { gt = 1 } } } 79 if gt == 1 { bestv = logits[i]; best = i } 80 i = i + 1 81 } 82 let tb: *i64 = sys_mmap(8) as *i64 83 tb[0] = best as i64 84 let db: *u8 = sys_mmap(64) 85 let dn: nx_int = nx_bpe_decode_bytelevel(vocab, tb, 1, db) 86 pr_puts(" #"); pr_num((r+1) as i64); pr_puts(" tok="); pr_num(best as i64); pr_puts(" [") 87 sys_write(1, db, dn as i64); pr_puts("]\n" as *u8) 88 logits[best] = NEG 89 r = r + 1 90 } 91 pr_puts("(top-1 = ' Paris' or coherent => FORWARD CORRECT, loop was the sampler; garbage => forward residual)\n" as *u8) 92 return 0 93}