code wiki / (root) / nx_f32_llm_run.nx

nx_f32_llm_run.nx source

↩ module page · 155 lines · 6090 B

1// nx_f32_llm_run.nx -- autoregressive generation loop runner. 2// 3// The capstone API surface: takes a populated model + vocab + cache, 4// a prompt as raw bytes, and produces generated bytes via prefill + 5// repeated (sample, decode-forward) loop. 6// 7// Composes every prior brick in the f32 LLM pipeline: 8// nx_bpe_encode prompt bytes -> token ids 9// nx_f32_llm_forward prefill + per-token decode forward 10// nx_f32_sampler_sample_top_k (or argmax if top_k <= 0) 11// nx_bpe_decode token id -> output bytes 12// nx_f32_kv_cache growing cache across forward calls 13// 14// API: 15// nx_f32_llm_run(model, vocab, cache, prompt, n_prompt, 16// max_new_tokens, inv_temp_f32, top_k, 17// eps, attn_scale, rope_log_base, apply_rope, 18// prng_state, eos_token_id, out_bytes, 19// out_bytes_cap) -> nx_int (= n bytes emitted) 20// 21// 16 args -- at the NishiLang argument limit. 22// 23// Termination: 24// * eos_token_id matched -> stop (caller passes -1 to disable) 25// * max_new_tokens reached -> stop 26// * out_bytes_cap reached -> stop 27// 28// Returns the count of bytes emitted into out_bytes (>= 0), or a 29// negative error verdict. 30// 31// genealogy_id: standard_autoregressive_decode_loop 32// lineage_id: substrate_f32_llm_run_v1 33 34import "nx_syscalls.nx" 35import "nx_tier.nx" 36import "nx_bpe.nx" 37import "nx_f32_kv_cache.nx" 38import "nx_f32_llm.nx" 39import "nx_f32_sampler.nx" 40 41const NX_FRN_OK_BASE: nx_int = 0 // success: emitted bytes count returned 42const NX_FRN_ERR_NULL: nx_int = -1 43const NX_FRN_ERR_BAD_DIM: nx_int = -2 44const NX_FRN_ERR_FORWARD: nx_int = -3 45const NX_FRN_ERR_TOK_OVF: nx_int = -4 46 47// Helper: pick next token id based on sampler config. 48 49func nx_f32_llm_run_pick(logits: *i64, vocab_size: nx_int, 50 top_k: nx_int, inv_temp_f32: i64, 51 prng_state: *i64) -> nx_int { 52 if top_k <= 0 { 53 return nx_f32_sampler_argmax(logits, vocab_size) 54 } 55 if prng_state == (0 as *i64) { 56 return nx_f32_sampler_argmax(logits, vocab_size) 57 } 58 return nx_f32_sampler_sample_top_k(logits, vocab_size, top_k, 59 inv_temp_f32, prng_state) 60} 61 62func nx_f32_llm_run(model: *NxF32LlamaModel, 63 vocab: *NxBpeVocab, 64 cache: *NxF32KVCache, 65 prompt: *u8, 66 n_prompt: nx_int, 67 max_new_tokens: nx_int, 68 inv_temp_f32: i64, 69 top_k: nx_int, 70 eps: i64, 71 attn_scale: i64, 72 rope_log_base: i64, 73 apply_rope: nx_int, 74 prng_state: *i64, 75 eos_token_id: nx_int, 76 out_bytes: *u8, 77 out_bytes_cap: nx_int) -> nx_int { 78 79 if model == (0 as *NxF32LlamaModel) { return NX_FRN_ERR_NULL } 80 if vocab == (0 as *NxBpeVocab) { return NX_FRN_ERR_NULL } 81 if cache == (0 as *NxF32KVCache) { return NX_FRN_ERR_NULL } 82 if out_bytes_cap <= 0 { return NX_FRN_ERR_BAD_DIM } 83 if max_new_tokens < 0 { return NX_FRN_ERR_BAD_DIM } 84 85 // ===== 1. Tokenize prompt ===== 86 // Worst case: 1 token per byte (no merges apply). 87 let prompt_tokens: *i64 = sys_mmap(n_prompt * 8) as *i64 88 let n_prompt_tok: nx_int = nx_bpe_encode(vocab, prompt, n_prompt, prompt_tokens) 89 if n_prompt_tok < 0 { return NX_FRN_ERR_BAD_DIM } 90 if n_prompt_tok == 0 { return NX_FRN_ERR_BAD_DIM } 91 92 // ===== 2. Allocate logits scratch ===== 93 let logits: *i64 = sys_mmap(n_prompt_tok * model.vocab_size * 8) as *i64 94 let logits1: *i64 = sys_mmap(model.vocab_size * 8) as *i64 95 96 // ===== 3. Prefill ===== 97 let v_pre: nx_int = nx_f32_llm_forward(model, prompt_tokens, n_prompt_tok, 98 cache, eps, attn_scale, rope_log_base, 99 apply_rope, logits) 100 if v_pre != NX_F32_LLM_OK { return NX_FRN_ERR_FORWARD } 101 102 // Get last-token logits row from prefill output. 103 let last_row: *i64 = (((logits as i64) + 104 (n_prompt_tok - 1) * model.vocab_size * 8)) as *i64 105 106 var n_emitted: nx_int = 0 107 let next_buf: *i64 = sys_mmap(8) as *i64 108 let one_byte_buf: *u8 = sys_mmap(64) 109 110 // ===== 4. Sample + decode loop ===== 111 var step: nx_int = 0 112 var current_logits: *i64 = last_row 113 while step < max_new_tokens { 114 // Sample. 115 let next_id: nx_int = nx_f32_llm_run_pick(current_logits, model.vocab_size, 116 top_k, inv_temp_f32, prng_state) 117 118 // EOS check. 119 if eos_token_id >= 0 { 120 if next_id == eos_token_id { 121 step = max_new_tokens // break 122 } 123 } 124 if step < max_new_tokens { 125 // Detokenize. 126 next_buf[0] = next_id as i64 127 let nb: nx_int = nx_bpe_decode(vocab, next_buf, 1, one_byte_buf) 128 if nb > 0 { 129 var bi: nx_int = 0 130 while bi < nb { 131 if n_emitted >= out_bytes_cap { 132 bi = nb 133 step = max_new_tokens // double break 134 } else { 135 out_bytes[n_emitted] = one_byte_buf[bi] 136 n_emitted = n_emitted + 1 137 bi = bi + 1 138 } 139 } 140 } 141 142 if step < max_new_tokens { 143 // Forward decode 1 token. 144 let v_d: nx_int = nx_f32_llm_forward(model, next_buf, 1, cache, 145 eps, attn_scale, rope_log_base, 146 apply_rope, logits1) 147 if v_d != NX_F32_LLM_OK { return NX_FRN_ERR_FORWARD } 148 current_logits = logits1 149 } 150 step = step + 1 151 } 152 } 153 154 return n_emitted 155}