Skip to main content

cherenkov/qwen4_exp/gpu/
decode.rs

1//! Trunk row-batch execution and single-token stepping.
2
3use super::*;
4
5impl Gpu<'_> {
6    /// Run `tokens` at positions pos.. as one batch; returns the greedy
7    /// argmax after each row. Nothing is committed until `commit`.
8    /// `snap` keeps rollback snapshots after each row but the last.
9    /// With `fold_mtp`, the MTP head's first pass runs in the same
10    /// command buffer over all `nb` rows, pairing row b with the trunk's
11    /// own prediction for it (the token that follows any accepted row),
12    /// so `mtp_draft` needs no round trip for the first draft.
13    pub fn step_rows(&mut self, tokens: &[u32], snap: bool, fold_mtp: bool) -> Result<Vec<u32>> {
14        let nb = tokens.len();
15
16        anyhow::ensure!(
17            !fold_mtp || self.mtp.is_some(),
18            "folding the MTP pass needs the MTP head"
19        );
20
21        self.folded_mtp = None;
22
23        anyhow::ensure!(
24            (1..=MAX_NB).contains(&nb),
25            "rows per step must be 1..={MAX_NB}"
26        );
27        anyhow::ensure!(self.pos + nb <= self.max_t, "context capacity exceeded");
28        anyhow::ensure!(
29            !snap || nb - 1 <= MAX_SNAP,
30            "at most {MAX_SNAP} draft rows per verify step"
31        );
32
33        let t0 = std::time::Instant::now();
34        let ngram0 = self.ngram_gather_s.get();
35        let c = &self.p.cfg;
36        let h = c.hidden_size as u32;
37        let hh = c.hc_hidden();
38
39        self.tokens.truncate(self.pos);
40        self.tokens.extend_from_slice(tokens);
41        self.ngram_prefetch_start(self.pos, nb);
42
43        self.batch_pos = self.pos;
44        self.batch_nb = nb;
45        self.batch_snap = snap;
46
47        unsafe {
48            let ids = self.scratch.ids.contents().cast::<u32>().as_ptr();
49
50            for (i, &t) in tokens.iter().enumerate() {
51                ids.add(IDS_IN + i).write(t);
52            }
53        }
54
55        self.last_experts.clear();
56        self.last_routes.clear();
57        self.last_states.clear();
58        self.last_route_w.clear();
59        self.last_miss.clear();
60        self.dispatch_count.set(0);
61
62        self.step_no += 1;
63        self.step_misses = 0;
64        self.step_miss_bytes = 0;
65        self.step_set_s = 0.0;
66        self.step_read_s = 0.0;
67        self.step_warm = 0;
68        self.step_cut = 0;
69        let snap_after = if snap { nb - 1 } else { 0 };
70        let pos = self.pos;
71        let n_layers = self.layers.len().min(layer_cap());
72        let base = self.event_base;
73        self.event_base += 4 * (n_layers as u64 + 1);
74
75        let phase_clock = self.phase_clock();
76        let cb = self.ctx.queue.commandBuffer().context("command buffer")?;
77        let mut enc = self.phase_encoder(&cb, None, 0)?;
78
79        {
80            let s = &self.scratch;
81            let (i0, nbu) = (IDS_IN as u32, nb as u32);
82
83            self.dispatch(
84                &enc,
85                &self.pipes.embed_rows,
86                |e| {
87                    self.bind(e, 0, &s.ids, 0);
88                    self.bind(e, 1, &self.dense, self.embed.w);
89                    self.bind(e, 2, &self.dense, self.embed.s);
90                    self.bind(e, 3, &self.dense, self.embed.b);
91                    self.bind(e, 4, &s.e, 0);
92                    set_bytes(e, 5, &h);
93                    set_bytes(e, 6, &i0);
94                    set_bytes(e, 7, &nbu);
95                },
96                nb * h as usize,
97                256,
98                false,
99            );
100
101            let gp = self.group_params(false, 0.0);
102
103            self.dispatch(
104                &enc,
105                &self.pipes.replicate_b,
106                |e| {
107                    self.bind(e, 0, &s.e, 0);
108                    self.bind(e, 1, &s.hyper, 0);
109                    set_bytes(e, 2, &gp);
110                    set_bytes(e, 3, &nbu);
111                },
112                nb * hh,
113                256,
114                false,
115            );
116        }
117
118        // A block's MoE output is injected by the next fused norm
119        // (`pending`); only the PLE block needs it applied up front.
120        let mut pending: Option<&Buf> = None;
121        let mtp_fold = if fold_mtp { self.mtp.as_ref() } else { None };
122
123        for li in 0..n_layers {
124            let layer = &self.layers[li];
125
126            self.encode_ple_before_block(&enc, layer, nb, pos, &mut pending)?;
127
128            // The last trunk layer looks ahead into the MTP block (its
129            // stream differs by the fold, an approximation like the rest).
130            let la = self.lookahead_layer(li, n_layers, fold_mtp);
131
132            self.encode_block(
133                &cb,
134                &mut enc,
135                layer,
136                li,
137                pos,
138                nb,
139                snap_after,
140                &self.scratch.hyper,
141                pending.take(),
142                la,
143                base + 4 * li as u64 + 1,
144            )?;
145
146            pending = Some(&self.scratch.moe_out);
147        }
148
149        self.head_b(
150            &enc,
151            &self.final_mixer,
152            nb,
153            &self.scratch.hyper,
154            0,
155            pending.take(),
156            &self.scratch.logits,
157            IDS_OUT,
158        );
159
160        if let Some(mtp) = mtp_fold {
161            // Row b of the head pairs the trunk's residual at pos + b with
162            // the trunk's argmax for it (written by the head just above).
163            self.encode_mtp_prelude(&enc, mtp, nb, IDS_OUT, &self.scratch.hyper, 0);
164
165            let slot_row = self.layers.len();
166
167            self.encode_block(
168                &cb,
169                &mut enc,
170                &mtp.layer,
171                slot_row,
172                pos,
173                nb,
174                0,
175                &self.scratch.mtp_hyper,
176                None,
177                None,
178                base + 4 * n_layers as u64 + 1,
179            )?;
180            self.head_b(
181                &enc,
182                &mtp.mixer,
183                nb,
184                &self.scratch.mtp_hyper,
185                0,
186                Some(&self.scratch.moe_out),
187                &self.scratch.mtp_logits,
188                IDS_MTP_OUT,
189            );
190        }
191
192        enc.endEncoding();
193        cb.commit();
194
195        // Service the layers in order while the GPU runs.
196        let mut predicted: std::collections::VecDeque<(usize, Vec<u32>)> =
197            std::collections::VecDeque::new();
198        let (mut la_hits, mut la_total, mut la_issued) = (0usize, 0usize, 0usize);
199        let mut io_s = 0.0f64;
200        let mut turn_s = 0.0f64;
201        let mtp_record = if fold_mtp {
202            self.mtp.as_ref().map(|m| m.layer.moe.record_layer)
203        } else {
204            None
205        };
206
207        for li in 0..n_layers {
208            let record_layer = self.layers[li].moe.record_layer;
209            let next = self
210                .lookahead_layer(li, n_layers, fold_mtp)
211                .map(|layer| layer.moe.record_layer);
212            let (hits, total, issued) = self.service_block(
213                record_layer,
214                li,
215                nb,
216                base + 4 * li as u64 + 1,
217                &mut predicted,
218                next,
219                &mut io_s,
220                &mut turn_s,
221            )?;
222            la_hits += hits;
223            la_total += total;
224            la_issued += issued;
225        }
226
227        if let Some(record_layer) = mtp_record {
228            let slot_row = self.layers.len();
229            let (hits, total, _) = self.service_block(
230                record_layer,
231                slot_row,
232                nb,
233                base + 4 * n_layers as u64 + 1,
234                &mut predicted,
235                None,
236                &mut io_s,
237                &mut turn_s,
238            )?;
239            la_hits += hits;
240            la_total += total;
241        }
242
243        cb.waitUntilCompleted();
244
245        self.collect_trunk_phases(n_layers, mtp_record, phase_clock);
246
247        let gpu_s = cb.GPUEndTime() - cb.GPUStartTime();
248
249        self.join_pending()?;
250        self.join_inflight()?;
251        self.gpu_ms.push(gpu_s * 1e3);
252        self.gpu_idle_ms.push(turn_s * 1e3);
253        self.dispatches.push(self.dispatch_count.get());
254        self.io_ms.push(io_s * 1e3);
255        self.set_ms.push(self.step_set_s * 1e3);
256        self.warm.push(self.step_warm);
257        self.read_ms.push(self.step_read_s * 1e3);
258        self.misses.push(self.step_misses);
259        self.cut.push(self.step_cut);
260        self.miss_bytes.push(self.step_miss_bytes);
261        self.lookahead_hit.push(if la_total > 0 {
262            la_hits as f64 / la_total as f64
263        } else {
264            f64::NAN
265        });
266        self.lookahead_issued.push(la_issued);
267        self.expert_history.push(self.last_experts.clone());
268        self.route_history.push((
269            tokens.to_vec(),
270            std::mem::take(&mut self.last_routes),
271            std::mem::take(&mut self.last_route_w),
272            std::mem::take(&mut self.last_miss),
273        ));
274
275        if self.dump_states {
276            self.state_history
277                .push(std::mem::take(&mut self.last_states));
278        }
279
280        self.rows.push(nb);
281        self.step_ms.push(t0.elapsed().as_secs_f64() * 1e3);
282        self.ngram_ms
283            .push((self.ngram_gather_s.get() - ngram0) * 1e3);
284
285        if fold_mtp {
286            self.folded_mtp =
287                Some(self.read_u32(&self.scratch.ids, IDS_MTP_OUT + nb)[IDS_MTP_OUT..].to_vec());
288        }
289
290        Ok(self.read_u32(&self.scratch.ids, IDS_OUT + nb)[IDS_OUT..].to_vec())
291    }
292
293    /// PLE needs the pending expert output injected before its convolution.
294    fn encode_ple_before_block(
295        &self,
296        enc: &Enc,
297        layer: &GLayer,
298        nb: usize,
299        pos: usize,
300        pending: &mut Option<&Buf>,
301    ) -> Result<()> {
302        let Some(pl) = &layer.ple else {
303            return Ok(());
304        };
305
306        if let Some(out) = pending.take() {
307            self.inject_b(enc, &self.scratch.hyper, out, nb);
308        }
309
310        self.ple_b(enc, pl, nb, pos)
311    }
312
313    /// One-block lookahead; deeper routing raised misses about 40%.
314    /// The last trunk layer predicts the folded MTP block when enabled.
315    fn lookahead_layer(&self, index: usize, count: usize, fold_mtp: bool) -> Option<&GLayer> {
316        if !self.lookahead {
317            return None;
318        }
319
320        if index + 1 < count {
321            return Some(&self.layers[index + 1]);
322        }
323
324        if fold_mtp {
325            return self.mtp.as_ref().map(|mtp| &mtp.layer);
326        }
327
328        None
329    }
330
331    /// One token, committed: the greedy next token.
332    pub fn step(&mut self, token: u32) -> Result<u32> {
333        let out = self.step_rows(&[token], false, false)?;
334
335        self.commit(1)?;
336
337        Ok(out[0])
338    }
339}