cherenkov/runner/
decode.rs1use super::*;
4use crate::sampling::Sampler;
5
6pub(crate) struct Decode {
7 cur: u32,
8 drafts: Vec<u32>,
9 n_draft: usize,
10 sampler: Option<Sampler>,
11 eos: Option<[u32; 2]>,
12 limit: usize,
13 pub tokens: Vec<u32>,
14 pub finish_reason: Option<&'static str>,
15 pub steps: usize,
16 pub accepted: usize,
17}
18
19#[cfg(test)]
20#[path = "../../tests/unit/runner/decode.rs"]
21mod tests;
22
23impl Decode {
24 pub(crate) fn rng(&self) -> Option<rand::rngs::StdRng> {
25 self.sampler.as_ref().map(Sampler::rng)
26 }
27
28 pub(crate) fn new(
29 seed: PrefillResume,
30 ids: &[u32],
31 tok: &tok::ChatTokenizer,
32 options: &Options,
33 limit: usize,
34 rng: Option<rand::rngs::StdRng>,
35 ) -> Result<Self> {
36 let mut sampler = if options.sampling.greedy() {
37 None
38 } else {
39 Some(Sampler::new(
40 &options.sampling,
41 ids,
42 seed.logits.as_ref().map_or(0, Vec::len),
43 )?)
44 };
45
46 if let (Some(sampler), Some(rng)) = (&mut sampler, rng) {
47 sampler.set_rng(rng);
48 }
49
50 let cur = match &mut sampler {
51 Some(s) => s.sample(
52 seed.logits
53 .as_deref()
54 .ok_or_else(|| anyhow::anyhow!("sampling needs final prompt logits"))?,
55 )?,
56 None => seed.next,
57 };
58
59 Ok(Self {
60 cur,
61 drafts: seed.drafts,
62 n_draft: options.effective_drafts(),
63 sampler,
64 eos: (!options.no_eos).then_some([tok.im_end, tok.endoftext]),
65 limit,
66 tokens: Vec::new(),
67 finish_reason: None,
68 steps: 0,
69 accepted: 0,
70 })
71 }
72
73 fn is_eos(&self, token: u32) -> bool {
74 self.eos.is_some_and(|ids| ids.contains(&token))
75 }
76
77 fn emit(&mut self, token: u32, emit: &mut dyn FnMut(u32) -> Result<()>) -> Result<()> {
78 self.tokens.push(token);
79
80 if let Some(sampler) = &mut self.sampler {
81 sampler.accept(token)?;
82 }
83
84 emit(token)
85 }
86
87 pub(crate) fn step(
89 &mut self,
90 gpu: &mut qwen4_exp::gpu::Gpu<'_>,
91 cpu: Option<&qwen4_exp::cpu::CpuModel<'_>>,
92 cpu_state: Option<&mut qwen4_exp::cpu::State>,
93 emit: &mut dyn FnMut(u32) -> Result<()>,
94 ) -> Result<()> {
95 if self.finish_reason.is_some() {
96 return Ok(());
97 }
98
99 if self.is_eos(self.cur) {
100 self.finish_reason = Some("stop");
101
102 return Ok(());
103 }
104
105 if self.tokens.len() >= self.limit {
106 self.finish_reason = Some("length");
107
108 return Ok(());
109 }
110
111 let mut rows = vec![self.cur];
112 let remaining = self.limit - self.tokens.len() - 1;
113
114 rows.extend(self.drafts.iter().take(self.n_draft.min(remaining)));
115 self.emit(self.cur, emit)?;
116
117 let pos = gpu.pos;
118 let result = gpu.step_rows(&rows, rows.len() > 1, self.n_draft > 0)?;
119 let mut n = 1;
120
121 while n < rows.len() && result[n - 1] == rows[n] && !self.is_eos(rows[n]) {
122 n += 1;
123 }
124
125 gpu.commit(n)?;
126
127 self.steps += 1;
128 self.accepted += n - 1;
129 let mut next = rows[1..n].to_vec();
130
131 next.push(result[n - 1]);
132
133 if self.n_draft > 0 {
134 let chain = if self.n_draft >= 2 && n < rows.len() {
135 1
136 } else {
137 self.n_draft
138 };
139 self.drafts = gpu.mtp_draft(&next, chain)?;
140 }
141
142 if let (Some(model), Some(state)) = (cpu, cpu_state) {
143 qwen4_exp_check_rows(
144 model,
145 state,
146 gpu,
147 &rows[..n],
148 (self.n_draft > 0).then_some(&next[..]),
149 pos,
150 "decode",
151 false,
152 )?;
153 }
154
155 for &token in &rows[1..n] {
156 self.emit(token, emit)?;
157 }
158
159 if self.tokens.len() == self.limit {
160 self.finish_reason = Some("length");
161
162 return Ok(());
163 }
164
165 self.cur = match &mut self.sampler {
166 Some(sampler) => sampler.sample(gpu.logits_row(n - 1))?,
167 None => result[n - 1],
168 };
169
170 Ok(())
171 }
172}