Skip to main content

cherenkov/qwen4_exp/gpu/prefill/
experts.rs

1//! Prefill expert grouping and event-protected streaming ring.
2
3use super::super::activity::reads::ReadSource;
4use super::*;
5
6// Pool entries stay resident for decode; ring entries are reused
7// only after the GPU signals that their previous batch is finished.
8enum ExpertSource {
9    Pool {
10        rid: usize,
11        buf: Buf,
12        record_offset: usize,
13        fetch: bool,
14    },
15    Ring(u32),
16}
17
18struct ExpertJob {
19    expert: usize,
20    csr_offset: usize,
21    row_count: usize,
22    source: ExpertSource,
23    layout: crate::qwen4_exp::lowbit::Layout,
24}
25
26impl Gpu<'_> {
27    /// Group prompt rows by expert and reserve the resident/ring destinations.
28    /// This preserves expert order, stable usage ranking and CSR row order.
29    fn prepare_expert_jobs(
30        &mut self,
31        record_layer: usize,
32        t: usize,
33        pf: &PrefillScratch,
34    ) -> Result<Vec<ExpertJob>> {
35        let c = &self.p.cfg;
36        let k = c.num_experts_per_tok;
37        let base_layout = crate::qwen4_exp::lowbit::Layout::four_bit(&self.p.manifest.experts);
38        let miss_layout = self.low_bit_store.unwrap_or(base_layout);
39        let n_record_layers = self.p.manifest.experts.layers;
40        let idx = self.read_u32(&pf.topk_idx, t * k);
41        let wts = self.read_f32(&pf.topk_w, t * k);
42        // Rows per expert, in expert order.
43        let mut lists: Vec<Vec<(u32, f32)>> = vec![Vec::new(); c.num_experts];
44
45        for r in 0..t {
46            for j in 0..k {
47                lists[idx[r * k + j] as usize].push((r as u32, wts[r * k + j]));
48            }
49        }
50
51        // The set keeps this layer's most-used experts (its share of the
52        // budget); everything else streams through the ring.
53        let budget = self.res.budget() / n_record_layers;
54        let mut by_use: Vec<usize> = (0..c.num_experts)
55            .filter(|&e| !lists[e].is_empty())
56            .collect();
57
58        by_use.sort_by_key(|&e| std::cmp::Reverse(lists[e].len()));
59
60        let keep: std::collections::HashSet<usize> = by_use.iter().take(budget).copied().collect();
61        let mut csr_rows: Vec<u32> = Vec::with_capacity(t * k);
62        let mut csr_w: Vec<f32> = Vec::with_capacity(t * k);
63        let mut jobs: Vec<ExpertJob> = Vec::new();
64        self.step_no += 1;
65        let mut ring_pos = 0usize;
66
67        for (e, list) in lists.iter().enumerate() {
68            if list.is_empty() {
69                continue;
70            }
71
72            let csr_offset = csr_rows.len();
73
74            for &(r, w) in list {
75                csr_rows.push(r);
76                csr_w.push(w);
77            }
78
79            let rid = self.record_id(record_layer, e as u32);
80
81            self.activity
82                .lookup(rid, list.len(), self.res.is_member(rid));
83
84            // Preserve a resident record's precision. New kept records
85            // use the pool's default; transient misses use --miss-experts.
86            let source = if keep.contains(&e) || self.res.is_member(rid) {
87                let fetch = !self
88                    .res
89                    .acquire(&self.ctx, &[rid], self.step_no)?
90                    .is_empty();
91                let (buf, record_offset) = self.res.buf(&self.ctx, rid)?;
92
93                ExpertSource::Pool {
94                    rid,
95                    buf,
96                    record_offset,
97                    fetch,
98                }
99            } else {
100                let slot = (ring_pos % RING) as u32;
101                ring_pos += 1;
102
103                ExpertSource::Ring(slot)
104            };
105            let layout = match &source {
106                ExpertSource::Pool { .. } if self.res.kind(rid) == 0 => base_layout,
107                _ => miss_layout,
108            };
109
110            let stats = &mut self.activity.layers[record_layer].quant[layout.kind() as usize];
111            stats.selected_experts += 1;
112            stats.selected_rows += list.len() as u64;
113
114            jobs.push(ExpertJob {
115                layout,
116                expert: e,
117                csr_offset,
118                row_count: list.len(),
119                source,
120            });
121        }
122
123        unsafe {
124            std::ptr::copy_nonoverlapping(
125                csr_rows.as_ptr(),
126                pf.csr_rows.contents().cast::<u32>().as_ptr(),
127                csr_rows.len(),
128            );
129            std::ptr::copy_nonoverlapping(
130                csr_w.as_ptr(),
131                pf.csr_w.contents().cast::<f32>().as_ptr(),
132                csr_w.len(),
133            );
134        }
135
136        Ok(jobs)
137    }
138
139    /// The MoE of one block over t rows: shared expert over all rows, then
140    /// the routed experts one at a time (their tokens gathered, computed,
141    /// scatter-added). The layer's most-used experts (up to the pool's
142    /// per-layer share) go into the residency set and stay for decode;
143    /// the rest stream through the ring. Returns (records fetched, bytes fetched, seconds
144    /// waiting on ring slots, GPU seconds).
145    pub(super) fn pf_experts(
146        &mut self,
147        moe: MoeRef,
148        t: usize,
149        pf: &PrefillScratch,
150    ) -> Result<(usize, usize, f64, f64)> {
151        let h = self.p.cfg.hidden_size as u32;
152        let inter = self.p.manifest.experts.inter as u32;
153        let stride = self.p.manifest.experts.record_stride as usize;
154        let base_layout = crate::qwen4_exp::lowbit::Layout::four_bit(&self.p.manifest.experts);
155        let miss_layout = self.low_bit_store.unwrap_or(base_layout);
156        let jobs = self.prepare_expert_jobs(moe.record_layer, t, pf)?;
157        let n_batches = jobs.len().div_ceil(GROUP);
158        // GPU -> CPU "batch done" on `event`, CPU -> GPU "batch fetched"
159        // on `event_cpu`; each has one writer, so values only grow.
160        let base = self.event_base;
161        self.event_base += n_batches as u64 + 1;
162        let cbase = self.event_cpu_base;
163        self.event_cpu_base += n_batches as u64 + 1;
164        let event: &ProtocolObject<dyn objc2_metal::MTLEvent> =
165            ProtocolObject::from_ref(&*self.event);
166        let event_cpu: &ProtocolObject<dyn objc2_metal::MTLEvent> =
167            ProtocolObject::from_ref(&*self.event_cpu);
168
169        let cb = self.ctx.queue.commandBuffer().context("command buffer")?;
170        let mut enc = cb.computeCommandEncoder().context("encoder")?;
171
172        // Shared expert over all rows, into a zeroed block output.
173        self.zero(&enc, &pf.moe_out, t as u32 * h);
174
175        if !self.skips("shared") {
176            self.qmm(&enc, &moe.sg, &pf.mixed, &pf.ge, t);
177            self.qmm(&enc, &moe.su, &pf.mixed, &pf.ue, t);
178            self.dispatch(
179                &enc,
180                &self.pipes.silu_mul,
181                |e| {
182                    self.bind(e, 0, &pf.ge, 0);
183                    self.bind(e, 1, &pf.ue, 0);
184                    self.bind(e, 2, &pf.hg, 0);
185                },
186                t * inter as usize,
187                256,
188                false,
189            );
190            self.qmm(&enc, &moe.sd, &pf.hg, &pf.ye, t);
191            self.dispatch(
192                &enc,
193                &self.pipes.shared_add_rows,
194                |e| {
195                    self.bind(e, 0, &pf.moe_out, 0);
196                    self.bind(e, 1, &pf.ye, 0);
197                    self.bind(e, 2, &self.dense, moe.gate.0);
198                    self.bind(e, 3, &pf.mixed, 0);
199                    set_bytes(e, 4, &h);
200                },
201                t,
202                256,
203                true,
204            );
205        }
206
207        if !self.skips("experts") {
208            for (bi, batch) in jobs.chunks(GROUP).enumerate() {
209                enc.endEncoding();
210                cb.encodeWaitForEvent_value(event_cpu, cbase + bi as u64 + 1);
211
212                enc = cb.computeCommandEncoder().context("encoder")?;
213
214                for job in batch {
215                    let (csr_offset, n) = (job.csr_offset, job.row_count);
216                    let nu = n as u32;
217                    let (wb, rec) = match &job.source {
218                        ExpertSource::Pool {
219                            buf, record_offset, ..
220                        } => (buf, *record_offset),
221                        ExpertSource::Ring(slot) => (&self.ring, *slot as usize * stride),
222                    };
223
224                    self.dispatch(
225                        &enc,
226                        &self.pipes.gather_rows,
227                        |e| {
228                            self.bind(e, 0, &pf.mixed, 0);
229                            self.bind(e, 1, &pf.csr_rows, csr_offset * 4);
230                            self.bind(e, 2, &pf.xg, 0);
231                            set_bytes(e, 3, &nu);
232                            set_bytes(e, 4, &h);
233                        },
234                        n * h as usize,
235                        256,
236                        false,
237                    );
238
239                    let l = job.layout;
240                    let gate = Q {
241                        w: l.gate_w,
242                        s: l.gate_s,
243                        b: l.gate_b,
244                        out: inter,
245                        inp: h,
246                    };
247                    let up = Q {
248                        w: l.up_w,
249                        s: l.up_s,
250                        b: l.up_b,
251                        out: inter,
252                        inp: h,
253                    };
254                    let down = Q {
255                        w: l.down_w,
256                        s: l.down_s,
257                        b: l.down_b,
258                        out: h,
259                        inp: inter,
260                    };
261
262                    self.expert_qmm_from(&enc, wb, &gate.at_offset(rec), &pf.xg, &pf.ge, n, l.bits);
263                    self.expert_qmm_from(&enc, wb, &up.at_offset(rec), &pf.xg, &pf.ue, n, l.bits);
264                    self.dispatch(
265                        &enc,
266                        &self.pipes.silu_mul,
267                        |e| {
268                            self.bind(e, 0, &pf.ge, 0);
269                            self.bind(e, 1, &pf.ue, 0);
270                            self.bind(e, 2, &pf.hg, 0);
271                        },
272                        n * inter as usize,
273                        256,
274                        false,
275                    );
276                    self.expert_qmm_from(&enc, wb, &down.at_offset(rec), &pf.hg, &pf.ye, n, l.bits);
277                    self.dispatch(
278                        &enc,
279                        &self.pipes.scatter_add_rows,
280                        |e| {
281                            self.bind(e, 0, &pf.ye, 0);
282                            self.bind(e, 1, &pf.csr_rows, csr_offset * 4);
283                            self.bind(e, 2, &pf.csr_w, csr_offset * 4);
284                            self.bind(e, 3, &pf.moe_out, 0);
285                            set_bytes(e, 4, &nu);
286                            set_bytes(e, 5, &h);
287                        },
288                        n * h as usize,
289                        256,
290                        false,
291                    );
292                }
293
294                enc.endEncoding();
295                cb.encodeSignalEvent_value(event, base + bi as u64 + 1);
296
297                enc = cb.computeCommandEncoder().context("encoder")?;
298            }
299        }
300
301        enc.endEncoding();
302        cb.commit();
303
304        // Stream the batches' records ahead of the GPU.
305        let mut fetched = 0usize;
306        let mut fetched_bytes = 0usize;
307        let mut wait_s = 0.0f64;
308        let ring_base = self.ring.contents().cast::<u8>().as_ptr() as usize;
309
310        for (bi, batch) in jobs.chunks(GROUP).enumerate() {
311            if bi >= RING / GROUP {
312                // The ring slots this batch overwrites were last used at
313                // most RING/GROUP batches ago; the GPU must be done there.
314                let need = base + (bi - RING / GROUP) as u64 + 1;
315                let t0 = std::time::Instant::now();
316
317                while self.event.signaledValue() < need {
318                    std::hint::spin_loop();
319                }
320
321                wait_s += t0.elapsed().as_secs_f64();
322            }
323
324            let (records, bytes) = self.read_expert_batch(
325                moe.record_layer,
326                batch,
327                ring_base,
328                stride,
329                miss_layout.kind(),
330            )?;
331            fetched += records;
332            fetched_bytes += bytes;
333
334            self.event_cpu.setSignaledValue(cbase + bi as u64 + 1);
335        }
336
337        cb.waitUntilCompleted();
338
339        Ok((
340            fetched,
341            fetched_bytes,
342            wait_s,
343            cb.GPUEndTime() - cb.GPUStartTime(),
344        ))
345    }
346
347    /// Read one batch into its reserved destinations before publishing its event.
348    fn read_expert_batch(
349        &mut self,
350        record_layer: usize,
351        batch: &[ExpertJob],
352        ring_base: usize,
353        stride: usize,
354        miss_kind: u8,
355    ) -> Result<(usize, usize)> {
356        let (mut to_set, mut ring4, mut ring_low) = (Vec::new(), Vec::new(), Vec::new());
357
358        for job in batch {
359            match &job.source {
360                ExpertSource::Pool {
361                    rid, fetch: true, ..
362                } => to_set.push(*rid),
363                ExpertSource::Ring(slot) => {
364                    let reads = if job.layout.bits == 4 {
365                        &mut ring4
366                    } else {
367                        &mut ring_low
368                    };
369
370                    reads.push(self.ring_read(
371                        record_layer,
372                        job,
373                        ring_base + *slot as usize * stride,
374                    ));
375                }
376                _ => {}
377            }
378        }
379
380        let fetched = to_set.len() + ring4.len() + ring_low.len();
381        let fetched_bytes = batch
382            .iter()
383            .filter(|job| !matches!(job.source, ExpertSource::Pool { fetch: false, .. }))
384            .map(|job| job.layout.stride)
385            .sum::<usize>();
386
387        if !self.fake_experts {
388            let (mut plan, _) = self.res.plan_reads(&to_set);
389
390            self.record_read_plan(&mut plan, &to_set, ReadSource::Prefill);
391
392            plan.run(&self.pool_file, &self.pool_file_nocache);
393            residency::fetch_into_slots(&self.pool_file_nocache, &ring4);
394
395            if !ring_low.is_empty() {
396                let file = self.res.store_file(&self.pool_file_nocache, miss_kind);
397
398                residency::fetch_into_slots(file, &ring_low);
399            }
400        }
401
402        self.res.finish(&self.ctx, &to_set)?;
403
404        Ok((fetched, fetched_bytes))
405    }
406
407    fn ring_read(
408        &mut self,
409        layer: usize,
410        job: &ExpertJob,
411        destination: usize,
412    ) -> residency::RecordRead {
413        let record = self.record_id(layer, job.expert as u32);
414        let mut read = residency::RecordRead {
415            destination,
416            file_offset: record * job.layout.stride,
417            bytes: job.layout.stride,
418            ticket: None,
419        };
420
421        if self.fake_experts {
422            return read;
423        }
424
425        self.activity.read(record, read.bytes);
426
427        read.ticket = Some(self.read_tracker.ticket(
428            layer,
429            job.layout.kind(),
430            ReadSource::Prefill,
431            read.bytes,
432        ));
433
434        read
435    }
436}