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}