1use super::activity::reads::{ReadSource, ReadTicket};
8use super::phases::CpuPhase;
9use super::*;
10
11pub(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
19struct BlockSignals {
22 router_ready: u64,
24 resident_ready: u64,
26 resident_done: u64,
28 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 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 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 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 fn wait_for_router(&self, ready: u64, slot_row: usize) -> Result<()> {
141 if !self.spin_wait {
142 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 #[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 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 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 #[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 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 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 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 #[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 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 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 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 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 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 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 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;