code wiki / (root) / nx_reason_paged_probe.nx

nx_reason_paged_probe.nx source

↩ module page · 126 lines · 4755 B

1// nx_reason_paged_probe.nx -- LIVE probe of the promoted organ entry 2// nx_reason_selfconsist_paged (prefix-shared SC): N=3 on the arithmetic 3// question; majority must be 85. expect_exit: 0 4 5import "nx_syscalls.nx" 6import "nx_tier.nx" 7import "nx_le.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" 24import "nx_f32_sampler.nx" 25import "nx_prng.nx" 26import "nx_reasoning.nx" 27import "nx_kvcache.nx" 28import "nx_f32_attn_paged.nx" 29import "nx_f32_llama_v4p.nx" 30import "nx_reasoning_paged.nx" 31 32func rp_w(s: *u8) -> i64 { var n: i64 = 0; while s[n] != (0 as u8) { n = n + 1 } sys_write(1, s, n); return 0 } 33func rp_wn(v: i64) -> i64 { 34 var m: i64 = v 35 if m < 0 { rp_w("-" as *u8); m = 0 - m } 36 let t: *u8 = sys_mmap(28) 37 var k: i64 = 0 38 if m == 0 { t[0] = 48 as u8; k = 1 } 39 while m > 0 { t[k] = (48 + (m % 10)) as u8; m = m / 10; k = k + 1 } 40 let o: *u8 = sys_mmap(28) 41 var i: i64 = 0 42 while i < k { o[i] = t[k - 1 - i]; i = i + 1 } 43 sys_write(1, o, k) 44 return 0 45} 46 47func main() -> i64 { 48 let path: *u8 = "/tmp/nx_real_model.gguf" as *u8 49 let len_out: *i64 = sys_mmap(8) as *i64 50 let buf: *u8 = sys_read_file(path, len_out) 51 if buf == (0 as *u8) { return 10 } 52 let hdr: *NxGgufHeader = sys_mmap(NX_GGUF_HDR_BYTES) as *NxGgufHeader 53 if nx_gguf_parse(buf, len_out[0], hdr) != NX_GGUF_OK { return 20 } 54 let model: *NxF32LlamaModel = nx_f32_llama_model_alloc() 55 let out_err: *i64 = sys_mmap(8) as *i64 56 if nx_f32_llm_read_dims_from_gguf(buf, len_out[0], hdr, model, out_err) != NX_FLD_OK { return 30 } 57 if nx_f32_llm_load_weights_v4_from_gguf(buf, hdr, model, out_err) != NX_FLV4_OK { return 40 } 58 let vocab: *NxBpeVocab = nx_bpe_vocab_new(67108864, 262144, 524288) 59 let nt2: *i64 = sys_mmap(8) as *i64 60 let nm: *i64 = sys_mmap(8) as *i64 61 if nx_f32_bpe_load_from_gguf(buf, len_out[0], hdr, vocab, nt2, nm, out_err) != NX_FBL_OK { return 50 } 62 let eos: nx_int = nx_f32_llm_read_eos(buf, len_out[0], hdr) 63 64 let rc: *NxReasonCfg = nx_reason_cfg_alloc() 65 rc.model = model 66 rc.vocab = vocab 67 rc.cache = 0 as *NxF32KVCache // unused on the paged path 68 rc.max_new = 8 69 rc.inv_temp_f32 = 0x3FA00000 70 rc.top_k = 40 71 rc.eps = 0x358637BD 72 rc.attn_scale = 0x3E000000 73 rc.rope_log_base = 0x415D0EAB 74 rc.eos = eos 75 rc.im_start = 151644 76 rc.im_end = 151645 77 78 let kv_dim: nx_int = model.n_kv_heads * model.head_dim 79 let pool: *NxPagedPool = nx_pkv_pool_new(16, model.n_layers, kv_dim) 80 81 let q: *u8 = "What is 47 plus 38? Answer with only the number." as *u8 82 let ans: *i64 = sys_mmap(3 * 8) as *i64 83 let lens: *i64 = sys_mmap(3 * 8) as *i64 84 let texts: *u8 = sys_mmap(3 * 96) 85 let t0: i64 = sys_now_us() 86 let mv: i64 = nx_reason_selfconsist_paged(rc, pool, q, 49, 3, 20260709, 87 ans, texts, 96, lens) 88 let us: i64 = sys_now_us() - t0 89 rp_w("PAGED-SC answers:" as *u8) 90 var i: nx_int = 0 91 while i < 3 { rp_w(" " as *u8); rp_wn(ans[i]); i = i + 1 } 92 rp_w(" -> majority=" as *u8); rp_wn(mv) 93 rp_w(" expected=85 us=" as *u8); rp_wn(us) 94 rp_w(" pool_free=" as *u8); rp_wn(pool.n_free) 95 rp_w("\n" as *u8) 96 if mv != 85 { return 60 } 97 if pool.n_free != 16 { return 61 } 98 99 // batched variant: same seeds -> identical answers, fewer forwards. 100 let ansB: *i64 = sys_mmap(3 * 8) as *i64 101 let lensB: *i64 = sys_mmap(3 * 8) as *i64 102 let textsB: *u8 = sys_mmap(3 * 96) 103 let t1: i64 = sys_now_us() 104 let mvB: i64 = nx_reason_selfconsist_batched(rc, pool, q, 49, 3, 20260709, 105 ansB, textsB, 96, lensB) 106 let usB: i64 = sys_now_us() - t1 107 rp_w("BATCHED-SC answers:" as *u8) 108 var i2: nx_int = 0 109 while i2 < 3 { rp_w(" " as *u8); rp_wn(ansB[i2]); i2 = i2 + 1 } 110 rp_w(" -> majority=" as *u8); rp_wn(mvB) 111 rp_w(" us=" as *u8); rp_wn(usB) 112 rp_w(" pool_free=" as *u8); rp_wn(pool.n_free) 113 rp_w("\n" as *u8) 114 if mvB != mv { return 62 } 115 var i3: nx_int = 0 116 while i3 < 3 { 117 if ansB[i3] != ans[i3] { return 63 } 118 if lensB[i3] != lens[i3] { return 64 } 119 i3 = i3 + 1 120 } 121 if pool.n_free != 16 { return 65 } 122 var us2: i64 = usB 123 if us2 < 1 { us2 = 1 } 124 rp_w("seq_vs_batched_x100=" as *u8); rp_wn(us * 100 / us2); rp_w("\n" as *u8) 125 return 0 126}