1use super::*;
4
5impl Gpu<'_> {
6 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 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 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 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 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 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 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 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}