Skip to main content

cherenkov/qwen4_exp/gpu/
streaming.rs

1//! Decode expert streaming: GPU routing, CPU slot publication, and file IO.
2//!
3//! The router and fetched-part signals share `event`; resident completion
4//! uses `event_res`. Event waits protect the CPU/GPU ownership transitions
5//! of the routing scratch, slot table, and expert records.
6
7use super::activity::reads::{ReadSource, ReadTicket};
8use super::phases::CpuPhase;
9use super::*;
10
11/// A background read may be absent under the FAKE developer option, but
12/// its reserved records must still pass through residency completion.
13pub(super) struct PendingRead {
14    pub(super) thread: Option<std::thread::JoinHandle<()>>,
15    pub(super) records: Vec<usize>,
16    pub(super) tickets: Vec<(usize, ReadTicket)>,
17}
18
19/// One block's monotonic handshake. The two event objects have distinct
20/// roles: do not publish `resident_done` on the router/table event.
21struct BlockSignals {
22    /// GPU -> CPU on event: routing scratch is readable.
23    router_ready: u64,
24    /// CPU -> GPU on event: table published; resident experts may run.
25    resident_ready: u64,
26    /// GPU -> CPU on event_res: optional weak-miss deadline boundary.
27    resident_done: u64,
28    /// CPU -> GPU on event: required misses are readable and resident.
29    misses_ready: u64,
30}
31
32impl BlockSignals {
33    fn new(seq: u64) -> Self {
34        Self {
35            router_ready: seq,
36            resident_ready: seq + 1,
37            resident_done: seq + 2,
38            misses_ready: seq + 3,
39        }
40    }
41}
42
43impl Gpu<'_> {
44    /// Wait for reads the deadline policy stopped waiting on, then make
45    /// their records usable (nothing reuses their slots before this).
46    pub(super) fn join_inflight(&mut self) -> Result<()> {
47        if self.inflight.is_empty() {
48            return Ok(());
49        }
50
51        let list = std::mem::take(&mut self.inflight);
52        let rids: Vec<usize> = list.iter().map(|(_, r)| *r).collect();
53
54        for (l, _) in &list {
55            l.wait()?;
56        }
57
58        self.res.finish(&self.ctx, &rids)
59    }
60
61    /// Join the background read of predicted records and add them to the
62    /// set (they are in memory once the read is done).
63    pub(super) fn join_pending(&mut self) -> Result<()> {
64        let Some(read) = self.pending.take() else {
65            return Ok(());
66        };
67        let t = std::time::Instant::now();
68
69        if let Some(thread) = read.thread {
70            let _ = thread.join();
71        }
72
73        self.step_read_s += t.elapsed().as_secs_f64();
74        let t = std::time::Instant::now();
75
76        self.res.finish(&self.ctx, &read.records)?;
77
78        self.step_set_s += t.elapsed().as_secs_f64();
79
80        Ok(())
81    }
82
83    /// Read predicted records for the next layer on a background thread
84    /// while the GPU runs the current layer's experts; they join the set
85    /// at the next service. Returns records read.
86    fn prefetch_async(&mut self, record_layer: usize, ids: &[u32]) -> Result<usize> {
87        self.join_pending()?;
88
89        let rids: Vec<usize> = ids
90            .iter()
91            .map(|&e| self.record_id(record_layer, e))
92            .collect();
93        let t_acq = std::time::Instant::now();
94
95        for &rid in &rids {
96            self.activity.records[rid].prefetch_requests += 1;
97            self.activity.layers[record_layer]
98                .prediction
99                .predicted_resident += u64::from(self.res.is_member(rid));
100        }
101
102        let need = self.res.acquire(&self.ctx, &rids, self.step_no)?;
103        self.step_set_s += t_acq.elapsed().as_secs_f64();
104
105        if need.is_empty() {
106            return Ok(0);
107        }
108
109        let (mut plan, warm) = self.res.plan_reads(&need);
110
111        self.record_read_plan(&mut plan, &need, ReadSource::Prefetch);
112
113        self.step_warm += warm;
114        let n = plan.len();
115
116        if self.fake_experts {
117            self.pending = Some(PendingRead {
118                thread: None,
119                tickets: Vec::new(),
120                records: need,
121            });
122
123            return Ok(n);
124        }
125
126        let (cached, nocache) = (
127            self.pool_file.try_clone().expect("dup experts fd"),
128            self.pool_file_nocache.try_clone().expect("dup experts fd"),
129        );
130        self.pending = Some(PendingRead {
131            tickets: plan.tickets(&need),
132            thread: Some(std::thread::spawn(move || plan.run(&cached, &nocache))),
133            records: need,
134        });
135
136        Ok(n)
137    }
138
139    /// Wait until the GPU releases the routing scratch for CPU reads.
140    fn wait_for_router(&self, ready: u64, slot_row: usize) -> Result<()> {
141        if !self.spin_wait {
142            // Blocking avoids heating a CPU core on the fanless target machine.
143            let ok = self.event.waitUntilSignaledValue_timeoutMS(ready, 30_000);
144
145            anyhow::ensure!(
146                ok,
147                "GPU did not reach the router of slot row {slot_row} within 30 s"
148            );
149
150            return Ok(());
151        }
152
153        let started = std::time::Instant::now();
154
155        while self.event.signaledValue() < ready {
156            std::hint::spin_loop();
157            anyhow::ensure!(
158                started.elapsed().as_secs() <= 30,
159                "GPU did not reach the router of slot row {slot_row} within 30 s"
160            );
161        }
162
163        Ok(())
164    }
165
166    fn log_lookahead(
167        &mut self,
168        record_layer: usize,
169        slot_row: usize,
170        experts: &[u32],
171        indices: &[u32],
172        weights: &[f32],
173    ) {
174        if !self.log_la {
175            return;
176        }
177
178        let k = self.p.cfg.num_experts_per_tok;
179
180        for &expert in experts {
181            let (mut weight, mut rank) = (0.0f32, u32::MAX);
182
183            for (i, &selected) in indices.iter().enumerate() {
184                if selected != expert {
185                    continue;
186                }
187
188                weight = weight.max(weights[i]);
189                rank = rank.min((i % k) as u32);
190            }
191
192            let resident = self.res.is_member(self.record_id(record_layer, expert));
193
194            self.la_pending.push(LaEntry {
195                step: self.step_no,
196                layer: slot_row,
197                expert,
198                weight,
199                rank,
200                resident,
201                hit: false,
202            });
203        }
204    }
205
206    /// One decoder block (attention or DeltaNet, then MoE) over `nb` rows
207    /// of `hyper`, with the event handshake around the expert dispatch.
208    /// The block's MoE output is left in scratch.moe_out for the next
209    /// fused norm to inject.
210    #[allow(clippy::too_many_arguments)]
211    pub(super) fn encode_block(
212        &self,
213        cb: &ProtocolObject<dyn MTLCommandBuffer>,
214        enc: &mut Retained<Enc>,
215        layer: &GLayer,
216        slot_row: usize,
217        base_pos: usize,
218        nb: usize,
219        snap_after: usize,
220        hyper: &Buf,
221        pending: Option<&Buf>,
222        lookahead: Option<&GLayer>,
223        seq: u64,
224    ) -> Result<()> {
225        let signals = BlockSignals::new(seq);
226        let s = &self.scratch;
227        let event: &ProtocolObject<dyn objc2_metal::MTLEvent> =
228            ProtocolObject::from_ref(&*self.event);
229
230        self.hc_read_b(enc, &layer.attn_hc, nb, hyper, 0, pending, &s.hc, true);
231
232        if !self.skips("mixer") {
233            match &layer.mix {
234                Mix::Attn(a) => self.attention_b(enc, a, base_pos, nb),
235                Mix::Delta(d) => self.deltanet_b(enc, d, nb, snap_after),
236            }
237        }
238
239        self.hc_read_b(
240            enc,
241            &layer.mlp_hc,
242            nb,
243            hyper,
244            0,
245            Some(&s.mix_out),
246            &s.hc,
247            true,
248        );
249        self.router_b(
250            enc,
251            &layer.moe,
252            nb,
253            &s.hc.mixed,
254            &s.router,
255            &s.topk_idx,
256            &s.topk_w,
257            self.p.cfg.num_experts_per_tok,
258        );
259
260        if let Some(next) = lookahead {
261            // Approximate the next layer's routing on the current stream.
262            self.hc_read_b(enc, &next.mlp_hc, nb, hyper, 0, None, &s.la, false);
263            self.router_b(
264                enc,
265                &next.moe,
266                nb,
267                &s.la.mixed,
268                &s.la_router,
269                &s.la_idx,
270                &s.la_w,
271                self.p.cfg.num_experts_per_tok,
272            );
273        }
274
275        enc.endEncoding();
276        cb.encodeSignalEvent_value(event, signals.router_ready);
277        cb.encodeWaitForEvent_value(event, signals.resident_ready);
278
279        *enc = self.phase_encoder(cb, Some(4 * slot_row + 1), 4 * slot_row + 2)?;
280
281        self.experts_b(enc, &layer.moe, slot_row, nb, 0);
282        enc.endEncoding();
283
284        let event_res: &ProtocolObject<dyn objc2_metal::MTLEvent> =
285            ProtocolObject::from_ref(&*self.event_res);
286
287        cb.encodeSignalEvent_value(event_res, signals.resident_done);
288        cb.encodeWaitForEvent_value(event, signals.misses_ready);
289
290        *enc = self.phase_encoder(cb, Some(4 * slot_row + 3), 4 * slot_row + 4)?;
291
292        self.experts_b(enc, &layer.moe, slot_row, nb, 1);
293
294        Ok(())
295    }
296
297    /// Fetched records use the miss precision; resident records keep their kind.
298    fn miss_record_bytes(&mut self, need: &[usize]) -> usize {
299        let Some(store) = self.low_bit_store else {
300            return self.p.manifest.experts.record_stride as usize;
301        };
302
303        if !self.all_low_bits {
304            let kind = if store.bits == 2 { 2 } else { 1 };
305
306            for &rid in need {
307                self.res.set_kind(rid, kind);
308            }
309        }
310
311        store.stride
312    }
313
314    /// Wait for required reads and truncate late weak records at the deadline.
315    /// The caller has published the resident table and has not released misses.
316    #[allow(clippy::too_many_arguments)]
317    fn wait_for_deadline(
318        &mut self,
319        slot_row: usize,
320        resident_done: u64,
321        flags: &[residency::Landed],
322        need: &[usize],
323        rids: &[usize],
324        order: &[usize],
325        n_res: usize,
326        wmax: &[f32],
327    ) -> Result<()> {
328        let t = std::time::Instant::now();
329        let landed = |f: &residency::Landed| f.done();
330        // Whichever comes first: every read landed (release at once, as
331        // before), or the GPU finished the resident part (cut the weak
332        // stragglers, wait for the strong ones).
333        let mut cut_now = false;
334
335        loop {
336            if flags.iter().all(landed) {
337                break;
338            }
339
340            if self.event_res.signaledValue() >= resident_done {
341                cut_now = true;
342
343                break;
344            }
345
346            std::hint::spin_loop();
347
348            if t.elapsed().as_secs() > 30 {
349                anyhow::bail!(
350                    "GPU did not finish the resident experts of slot row {slot_row} within 30 s"
351                );
352            }
353        }
354
355        let mut cut_kinds = [false; 3];
356
357        if cut_now {
358            let need_pos = |rid: usize| need.iter().position(|&r| r == rid).unwrap();
359            let mut active_count = order.len();
360            // Late entries sit at order[n_res..], strongest first.
361            let mut u = order.len();
362
363            while u > n_res {
364                let i = order[u - 1];
365                let f = &flags[need_pos(rids[i])];
366
367                if landed(f) {
368                    break;
369                }
370
371                if wmax[i] < self.cut_w {
372                    active_count -= 1;
373                    self.step_cut += 1;
374                    let record = rids[i];
375                    let layer = self.activity.layer_for_record(record);
376                    let kind = self.res.kind(record) as usize;
377                    self.activity.layers[layer].quant[kind].cut_experts += 1;
378                    cut_kinds[kind] = true;
379
380                    self.inflight.push((f.clone(), rids[i]));
381
382                    u -= 1;
383                } else {
384                    break;
385                }
386            }
387
388            for &i in &order[n_res..active_count] {
389                flags[need_pos(rids[i])].wait()?;
390            }
391
392            unsafe {
393                let tab = self
394                    .slot_tab
395                    .contents()
396                    .cast::<u64>()
397                    .as_ptr()
398                    .add(slot_row * SLOT_STRIDE);
399
400                tab.add(SLOT_STRIDE - 1).write(active_count as u64);
401            }
402        }
403
404        if let Some(&record) = rids.first() {
405            let layer = self.activity.layer_for_record(record);
406
407            for (stats, cut) in self.activity.layers[layer].quant.iter_mut().zip(cut_kinds) {
408                stats.cut_batches += u64::from(cut);
409            }
410        }
411
412        self.step_read_s += t.elapsed().as_secs_f64();
413
414        Ok(())
415    }
416
417    /// Cut reads stay in flight until step end; publish only completed records.
418    fn finish_miss_reads(
419        &mut self,
420        need: &[usize],
421        flags: &[residency::Landed],
422        deadline: bool,
423    ) -> Result<()> {
424        if !deadline {
425            return self.res.finish(&self.ctx, need);
426        }
427
428        let landed: Vec<usize> = need
429            .iter()
430            .enumerate()
431            .filter(|(j, _)| flags.is_empty() || flags[*j].done())
432            .map(|(_, &r)| r)
433            .collect();
434
435        self.res.finish(&self.ctx, &landed)
436    }
437
438    /// CPU side of one block's handshake: wait for the router, publish the
439    /// union of the rows' experts with the resident ones first. Read misses
440    /// while the GPU computes the resident part, then issue lookahead and
441    /// release the fetched part. Returns (lookahead hits, total experts,
442    /// records read for lookahead). See BlockSignals for the event contract.
443    #[allow(clippy::too_many_arguments)]
444    pub(super) fn service_block(
445        &mut self,
446        record_layer: usize,
447        slot_row: usize,
448        nb: usize,
449        seq: u64,
450        predicted: &mut std::collections::VecDeque<(usize, Vec<u32>)>,
451        lookahead: Option<usize>,
452        io_s: &mut f64,
453        turn_s: &mut f64,
454    ) -> Result<(usize, usize, usize)> {
455        let signals = BlockSignals::new(seq);
456        let k = self.p.cfg.num_experts_per_tok;
457
458        self.wait_for_router(signals.router_ready, slot_row)?;
459
460        let t0 = std::time::Instant::now();
461        let cpu0 = super::phases::thread_cpu_seconds();
462
463        self.phase_cpu_mark(slot_row, CpuPhase::Observed);
464
465        let idx = self.read_u32(&self.scratch.topk_idx, nb * k);
466        let wts = self.read_f32(&self.scratch.topk_w, nb * k);
467
468        if self.dump_states {
469            // The GPU is stalled on the union, so this block's router
470            // input is still in scratch.
471            let h = self.p.cfg.hidden_size;
472
473            self.last_states
474                .push((self.read_f32(&self.scratch.hc.mixed, h), idx[..k].to_vec()));
475        }
476
477        let union = union_of(&idx);
478        let (mut hits, mut total) = (0, 0);
479
480        if predicted.front().is_some_and(|(row, _)| *row == slot_row) {
481            let (_, pred) = predicted.pop_front().unwrap();
482
483            self.record_prediction(record_layer, &pred, &union);
484
485            hits = union.iter().filter(|e| pred.contains(e)).count();
486            total = union.len();
487        }
488
489        for mut e in self.la_pending.drain(..) {
490            e.hit = union.contains(&e.expert);
491
492            self.la_log.push(e);
493        }
494
495        // Observe prediction readiness before joining its reads.
496        let prefetch_read_before = self.step_read_s;
497
498        self.join_pending()?;
499
500        self.activity.layers[record_layer]
501            .phases
502            .prefetch_wait_seconds += self.step_read_s - prefetch_read_before;
503
504        self.record_routed_experts(record_layer, &idx, &union);
505
506        let rids: Vec<usize> = union
507            .iter()
508            .map(|&e| self.record_id(record_layer, e))
509            .collect();
510
511        let t_acq = std::time::Instant::now();
512        let need = self.res.acquire(&self.ctx, &rids, self.step_no)?;
513        self.step_set_s += t_acq.elapsed().as_secs_f64();
514        self.step_misses += need.len();
515        let miss_bytes = self.miss_record_bytes(&need);
516        self.step_miss_bytes += need.len() * miss_bytes;
517
518        self.record_quant_selection(record_layer, &idx, &union);
519
520        // Start reading the misses at once, on their own thread; the GPU
521        // runs the resident experts meanwhile.
522        let ti = std::time::Instant::now();
523        let (mut plan, warm) = self.res.plan_reads(&need);
524
525        self.record_read_plan(&mut plan, &need, ReadSource::Demand);
526
527        self.step_warm += warm;
528        let deadline = self.cut_w > 0.0 && !self.fake_experts;
529        // Per-record completion only for the deadline policy; otherwise
530        // one thread for the batch. Metal IO lost at this read fan-out.
531        let mut flags: Vec<residency::Landed> = Vec::new();
532        let miss_read = if plan.is_empty() || self.fake_experts {
533            None
534        } else if deadline {
535            flags = plan.run_tracked(&self.pool_file, &self.pool_file_nocache, need.len());
536
537            None
538        } else {
539            let (cached, nocache) = (
540                self.pool_file.try_clone().expect("dup experts fd"),
541                self.pool_file_nocache.try_clone().expect("dup experts fd"),
542            );
543
544            Some(std::thread::spawn(move || plan.run(&cached, &nocache)))
545        };
546        // Resident records first, then the ones being fetched, strongest
547        // first (so a deadline cut is a truncation of the weak end).
548        let missing: Vec<bool> = rids.iter().map(|r| need.contains(r)).collect();
549        let wmax: Vec<f32> = union
550            .iter()
551            .map(|&e| {
552                (0..nb * k)
553                    .filter(|&i| idx[i] == e)
554                    .map(|i| wts[i])
555                    .fold(0.0f32, f32::max)
556            })
557            .collect();
558        let mut order: Vec<usize> = (0..union.len()).collect();
559
560        order.sort_by(|&a, &b| {
561            missing[a].cmp(&missing[b]).then(
562                wmax[b]
563                    .partial_cmp(&wmax[a])
564                    .unwrap_or(std::cmp::Ordering::Equal),
565            )
566        });
567
568        let n_res = missing.iter().filter(|m| !**m).count();
569        let addrs: Vec<u64> = order
570            .iter()
571            .map(|&i| self.res.addr(&self.ctx, rids[i]))
572            .collect::<Result<_>>()?;
573
574        unsafe {
575            let tab = self
576                .slot_tab
577                .contents()
578                .cast::<u64>()
579                .as_ptr()
580                .add(slot_row * SLOT_STRIDE);
581
582            for (u, &a) in addrs.iter().enumerate() {
583                tab.add(u).write(a);
584            }
585
586            tab.add(SLOT_STRIDE - 2).write(n_res as u64);
587            tab.add(SLOT_STRIDE - 1).write(union.len() as u64);
588
589            let wm = self
590                .wmap
591                .contents()
592                .cast::<f32>()
593                .as_ptr()
594                .add(slot_row * MAX_NB * SLOT_STRIDE);
595
596            for b in 0..nb {
597                let row = wm.add(b * SLOT_STRIDE);
598
599                for (u, &i) in order.iter().enumerate() {
600                    let e = union[i];
601                    let w = (0..k)
602                        .find(|&j| idx[b * k + j] == e)
603                        .map_or(0.0, |j| wts[b * k + j]);
604
605                    row.add(u).write(w);
606                }
607            }
608        }
609
610        self.phase_cpu_mark(slot_row, CpuPhase::ResidentRelease);
611        self.event.setSignaledValue(signals.resident_ready);
612        self.record_cut_eligible(record_layer, &rids, &missing, &wmax);
613
614        let read_wait_before = self.step_read_s;
615
616        // Deadline: when the resident part is done, cut the weak misses
617        // that have not landed (from the weak end, as a truncation), wait
618        // for the rest.
619        if deadline && !flags.is_empty() {
620            self.wait_for_deadline(
621                slot_row,
622                signals.resident_done,
623                &flags,
624                &need,
625                &rids,
626                &order,
627                n_res,
628                &wmax,
629            )?;
630        }
631
632        // Required misses land before lookahead starts; the deadline policy
633        // may leave cut reads in flight. Deferring prefetch measured 11%
634        // faster by keeping that traffic off the critical reads.
635        let mut issued = 0;
636
637        {
638            let t = std::time::Instant::now();
639
640            if let Some(h) = miss_read {
641                let _ = h.join();
642            }
643
644            self.step_read_s += t.elapsed().as_secs_f64();
645        }
646
647        self.activity.layers[record_layer]
648            .phases
649            .demand_wait_seconds += self.step_read_s - read_wait_before;
650
651        if let Some(next_layer) = lookahead {
652            let lk = self.p.cfg.num_experts_per_tok;
653            let la_idx = self.read_u32(&self.scratch.la_idx, nb * lk);
654            let la_w = self.read_f32(&self.scratch.la_w, nb * lk);
655            let la = union_of(&la_idx);
656
657            self.log_lookahead(next_layer, slot_row + 1, &la, &la_idx, &la_w);
658
659            issued = self.prefetch_async(next_layer, &la)?;
660
661            predicted.push_back((slot_row + 1, la));
662        }
663
664        let t_fin = std::time::Instant::now();
665
666        self.finish_miss_reads(&need, &flags, deadline)?;
667
668        self.step_set_s += t_fin.elapsed().as_secs_f64();
669        *io_s += ti.elapsed().as_secs_f64();
670
671        self.phase_cpu_mark(slot_row, CpuPhase::FetchedRelease);
672        self.event.setSignaledValue(signals.misses_ready);
673
674        let phase = &mut self.activity.layers[record_layer].phases;
675        phase.service_windows += 1;
676        phase.service_wall_seconds += t0.elapsed().as_secs_f64();
677        phase.service_cpu_seconds += (super::phases::thread_cpu_seconds() - cpu0).max(0.0);
678
679        *turn_s += t0.elapsed().as_secs_f64();
680
681        self.last_experts
682            .push(order.iter().map(|&i| union[i]).collect());
683        self.last_routes.push(idx);
684        self.last_route_w.push(wts);
685        self.last_miss.push(
686            (0..union.len())
687                .filter(|&i| missing[i])
688                .map(|i| union[i])
689                .collect(),
690        );
691
692        Ok((hits, total, issued))
693    }
694}
695
696#[cfg(test)]
697#[path = "../../../tests/unit/qwen4_exp/gpu/streaming.rs"]
698mod tests;