1use super::super::activity::reads::ReadSource;
4use super::*;
5
6enum 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 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 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 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 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 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 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 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 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 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 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}