cortiq_engine/pool.rs
1//! Persistent worker pool for row-parallel matvecs.
2//!
3//! Threads are spawned once and spin-then-park between calls — vmfcore
4//! measured spawn-per-matvec at ~+27% decode cost versus a persistent
5//! pool. Parallelism is by disjoint row ranges, so results are
6//! bit-identical to the serial path (each row's dot product is computed
7//! the same way).
8//!
9//! Dispatch is a single shared job slot + atomic epoch (roadmap §3 P0):
10//! the caller publishes one pointer, bumps the epoch and JOINS THE WORK
11//! as the extra worker instead of blocking on a latch. The previous
12//! design allocated an `Arc<Latch>` and pushed a message into every
13//! worker's mpsc channel for every matvec (~200 dispatches/token) —
14//! with decode-grade matvecs that synchronization was its own budget.
15//! Workers spin for `CMF_POOL_SPIN` iterations before parking.
16//! Default 4000: at ~39 dispatches/token, park-immediately pays the
17//! unpark syscall on every worker for every dispatch — measured on an
18//! M4 (interleaved A/B, current epoch dispatch + parked-flag design):
19//! Qwen-0.5B q8 decode 101→115 tok/s, q4t 117→149, the 50M bench model
20//! 549→954 at spin=4000 vs spin=0. An early measurement that showed
21//! spinning LOSING (−25% on q8) predates the parked-flag skip and the
22//! multi-matrix dispatch cuts; it no longer reproduces. Over-spinning
23//! still hurts (200k: −15% vs 4k — spinners steal the caller's serial
24//! cycles), so the budget stays bounded. `CMF_POOL_SPIN=0` restores
25//! park-immediately for share-the-box serving.
26//!
27//! `CMF_THREADS` env: 0/1 = serial, N = worker count
28//! (default: available_parallelism − 1, capped at 8).
29
30use std::cell::UnsafeCell;
31use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
32use std::sync::Arc;
33
34/// A `*const dyn Fn` that may cross a thread boundary. Safety is
35/// provided by `Pool::run`: the caller blocks until every worker has
36/// finished, so the borrow outlives all uses.
37#[derive(Clone, Copy)]
38struct TaskPtr(*const (dyn Fn(usize, usize) + Sync));
39unsafe impl Send for TaskPtr {}
40
41struct Inner {
42 /// Bumped once per published job; workers watch it.
43 epoch: AtomicUsize,
44 /// Workers still running the current job (excludes the caller).
45 remaining: AtomicUsize,
46 /// The published job: closure pointer + total participant count.
47 /// Written by the caller BEFORE the epoch bump, read by workers
48 /// AFTER they observe the new epoch (acquire/release pairing).
49 slot: UnsafeCell<Option<(TaskPtr, usize)>>,
50 shutdown: AtomicBool,
51 /// Spin iterations before a worker parks (0 = park immediately).
52 spin_budget: usize,
53 /// Per-worker "I am parked" flags — lets the caller skip the unpark
54 /// syscall for workers that are still spinning.
55 parked: Box<[AtomicBool]>,
56}
57
58// SAFETY: `slot` is only written while no job is in flight (run()
59// returns after `remaining` hits 0) and only read after the epoch
60// publication that follows the write.
61unsafe impl Sync for Inner {}
62
63/// Process-wide dispatch counter (roadmap §3 P0 «измерения»): one tick
64/// per published job. `bench --json` reports dispatches/token from it.
65static DISPATCHES: AtomicUsize = AtomicUsize::new(0);
66
67/// Total pool jobs published since process start (all pools).
68pub fn dispatch_count() -> usize {
69 DISPATCHES.load(Ordering::Relaxed)
70}
71
72/// Persistent thread pool: shared job slot, epoch dispatch, caller
73/// participation.
74pub struct Pool {
75 inner: Arc<Inner>,
76 /// Thread handles for `unpark` (same order as `parked`).
77 threads: Vec<std::thread::Thread>,
78 joins: Vec<std::thread::JoinHandle<()>>,
79}
80
81fn spin_budget_from_env() -> usize {
82 std::env::var("CMF_POOL_SPIN")
83 .ok()
84 .and_then(|v| v.parse::<usize>().ok())
85 .unwrap_or(4000)
86}
87
88impl Pool {
89 pub fn new(n_workers: usize) -> Self {
90 Self::with_spin(n_workers, spin_budget_from_env())
91 }
92
93 /// Explicit spin budget (tests pin it without touching the env).
94 pub fn with_spin(n_workers: usize, spin_budget: usize) -> Self {
95 let inner = Arc::new(Inner {
96 epoch: AtomicUsize::new(0),
97 remaining: AtomicUsize::new(0),
98 slot: UnsafeCell::new(None),
99 shutdown: AtomicBool::new(false),
100 spin_budget,
101 parked: (0..n_workers).map(|_| AtomicBool::new(false)).collect(),
102 });
103 let mut joins = Vec::with_capacity(n_workers);
104 for w in 0..n_workers {
105 let inner = inner.clone();
106 let h = std::thread::Builder::new()
107 .name(format!("cmf-pool-{w}"))
108 .spawn(move || worker_loop(&inner, w))
109 .expect("spawn pool worker");
110 joins.push(h);
111 }
112 let threads = joins.iter().map(|h| h.thread().clone()).collect();
113 Self { inner, threads, joins }
114 }
115
116 /// Big-core count on heterogeneous ARM (big.LITTLE): the kernel
117 /// exposes per-core capacity on Android and most ARM Linux; efficiency
118 /// cores in the pool DRAG the big ones on our row-parallel jobs (the
119 /// same cliff llama.cpp hits at -t 10 on an M4: 163 → 112 tok/s).
120 /// None = capacities absent or homogeneous.
121 #[cfg(all(target_arch = "aarch64", any(target_os = "linux", target_os = "android")))]
122 fn big_cores() -> Option<usize> {
123 let mut caps: Vec<u64> = Vec::new();
124 for cpu in 0.. {
125 let path = format!("/sys/devices/system/cpu/cpu{cpu}/cpu_capacity");
126 match std::fs::read_to_string(&path) {
127 Ok(v) => caps.push(v.trim().parse().ok()?),
128 Err(_) => break,
129 }
130 }
131 Self::cores_from_capacities(&caps)
132 }
133
134 /// How many cores the pool should use, from the kernel's per-core
135 /// capacity values. Capacity folds µarch × clock into one number,
136 /// and the two need different treatment: cores of ANOTHER µarch
137 /// (A5xx efficiency cluster next to A7xx/X: capacity ratio ≥ ~2)
138 /// drag row-parallel work down and are excluded; cores of the SAME
139 /// µarch merely clock-binned (JLQ JR510: 8×A55 as 4×2.0 + 4×1.5 GHz,
140 /// ratio 1.33) pull their weight and must ALL be used. The 1.6
141 /// threshold splits the two regimes: on a Snapdragon 8-class part
142 /// it keeps X + A7xx mid cores and drops A5xx.
143 #[cfg_attr(not(all(target_arch = "aarch64", any(target_os = "linux", target_os = "android"))), allow(dead_code))]
144 fn cores_from_capacities(caps: &[u64]) -> Option<usize> {
145 let max = *caps.iter().max()?;
146 let min = *caps.iter().min()?;
147 if caps.len() < 2 || max == min {
148 return None;
149 }
150 Some(caps.iter().filter(|&&c| c * 8 >= max * 5).count())
151 }
152
153 #[cfg(not(all(target_arch = "aarch64", any(target_os = "linux", target_os = "android"))))]
154 fn big_cores() -> Option<usize> {
155 None
156 }
157
158 /// Pool sized from `CMF_THREADS` (see module docs). `None` = serial.
159 /// Without the env, heterogeneous ARM defaults to its BIG cores.
160 pub fn from_env() -> Option<Arc<Self>> {
161 let n = match std::env::var("CMF_THREADS") {
162 Ok(v) => v.parse::<usize>().unwrap_or(0),
163 Err(_) => match Self::big_cores() {
164 Some(big) => big,
165 None => {
166 let avail = std::thread::available_parallelism()
167 .map(|n| n.get())
168 .unwrap_or(1);
169 avail.saturating_sub(1).min(8)
170 }
171 },
172 };
173 if n <= 1 {
174 None
175 } else {
176 Some(Arc::new(Self::new(n)))
177 }
178 }
179
180 /// Spawned worker threads (the caller joins each job on top).
181 pub fn n_workers(&self) -> usize {
182 self.threads.len()
183 }
184
185 /// Run `f(row_start, row_end)` over `0..rows`, self-balancing.
186 ///
187 /// One dispatch, but workers pull row-ranges from a shared cursor
188 /// instead of each taking a fixed 1/n slice. On a heterogeneous CPU
189 /// (Apple Silicon: 4 P-cores + 6 E-cores here) a static split makes
190 /// every matvec end at the SLOWEST core's pace while the fast ones
191 /// idle at the barrier; pulling by grain lets a P-core take several
192 /// chunks for each one an E-core takes, so skew collapses to a
193 /// single grain. Row ranges stay disjoint and each row's dot is
194 /// computed exactly as in the serial path → bit-identical output.
195 pub fn run_rows(&self, rows: usize, f: &(dyn Fn(usize, usize) + Sync)) {
196 // Enough chunks to balance, large enough to keep the SDOT inner
197 // loop and the hardware prefetcher in their stride.
198 let grain = (rows / ((self.threads.len() + 1) * 8)).max(32);
199 let next = AtomicUsize::new(0);
200 self.run(&|_w, _n| loop {
201 let start = next.fetch_add(grain, Ordering::Relaxed);
202 if start >= rows {
203 break;
204 }
205 f(start, (start + grain).min(rows));
206 });
207 }
208
209 /// Multi-matrix job: one dispatch serves SEVERAL row spaces
210 /// (roadmap §3 P0 — «одна внешняя публикация job на слой»). Parts
211 /// are laid out back-to-back in a virtual row space and pulled by
212 /// grain from one shared cursor, so QKV or gate+up cost a single
213 /// barrier instead of one each. Each part's `f(start, end)` sees its
214 /// OWN row indices — per-row math and outputs are bit-identical to
215 /// separate `run_rows` calls.
216 pub fn run_many(&self, parts: &[(usize, &(dyn Fn(usize, usize) + Sync))]) {
217 let total: usize = parts.iter().map(|p| p.0).sum();
218 if total == 0 {
219 return;
220 }
221 let grain = (total / ((self.threads.len() + 1) * 8)).max(32);
222 let next = AtomicUsize::new(0);
223 self.run(&|_w, _n| loop {
224 let s = next.fetch_add(grain, Ordering::Relaxed);
225 if s >= total {
226 break;
227 }
228 let e = (s + grain).min(total);
229 let mut base = 0usize;
230 for &(rows, f) in parts {
231 let a = s.max(base);
232 let b = e.min(base + rows);
233 if a < b {
234 f(a - base, b - base);
235 }
236 base += rows;
237 if base >= e {
238 break;
239 }
240 }
241 });
242 }
243
244 /// Run `f(worker_idx, n_participants)` on every worker AND the
245 /// calling thread (`worker_idx = n_workers()` for the caller);
246 /// returns when all participants have finished.
247 pub fn run(&self, f: &(dyn Fn(usize, usize) + Sync)) {
248 DISPATCHES.fetch_add(1, Ordering::Relaxed);
249 let nw = self.threads.len();
250 let n = nw + 1; // caller participates
251 // SAFETY: the wait loop below blocks until every worker is done,
252 // so extending the borrow to 'static never outlives the call.
253 let ptr: *const (dyn Fn(usize, usize) + Sync) = f;
254 let ptr: *const (dyn Fn(usize, usize) + Sync + 'static) =
255 unsafe { std::mem::transmute(ptr) };
256 // SAFETY: no job in flight (previous run() drained `remaining`),
257 // so the slot is not being read.
258 unsafe { *self.inner.slot.get() = Some((TaskPtr(ptr), n)) };
259 self.inner.remaining.store(nw, Ordering::Relaxed);
260 self.inner.epoch.fetch_add(1, Ordering::SeqCst);
261 for (i, t) in self.threads.iter().enumerate() {
262 if self.inner.parked[i].load(Ordering::SeqCst) {
263 t.unpark();
264 }
265 }
266
267 // The caller's share — the barrier costs nothing while there is
268 // real work to do.
269 f(nw, n);
270
271 // Wait for the stragglers (bounded by one worker's chunk).
272 let mut spins = 0usize;
273 while self.inner.remaining.load(Ordering::Acquire) != 0 {
274 spins += 1;
275 if spins < 10_000 {
276 std::hint::spin_loop();
277 } else {
278 std::thread::yield_now();
279 }
280 }
281 }
282}
283
284impl Drop for Pool {
285 fn drop(&mut self) {
286 self.inner.shutdown.store(true, Ordering::SeqCst);
287 for t in &self.threads {
288 t.unpark();
289 }
290 for h in self.joins.drain(..) {
291 let _ = h.join();
292 }
293 }
294}
295
296fn worker_loop(inner: &Inner, idx: usize) {
297 // The pool is created at epoch 0; baseline MUST be 0, not a fresh
298 // epoch read — if the caller publishes a job before the OS actually
299 // starts this thread, reading the live epoch would adopt that job's
300 // epoch as "already seen", skip it, and deadlock the caller's wait.
301 let mut seen = 0usize;
302 loop {
303 // Wait for a new epoch: spin first (decode publishes the next
304 // matvec within microseconds), park only when idle for real.
305 let mut spins = 0usize;
306 loop {
307 let e = inner.epoch.load(Ordering::Acquire);
308 if e != seen {
309 seen = e;
310 break;
311 }
312 if inner.shutdown.load(Ordering::Relaxed) {
313 return;
314 }
315 if spins < inner.spin_budget {
316 spins += 1;
317 std::hint::spin_loop();
318 } else {
319 inner.parked[idx].store(true, Ordering::SeqCst);
320 // Re-check under SeqCst: the caller bumps the epoch
321 // BEFORE reading `parked`, so either it sees our flag
322 // (and unparks) or we see its epoch here — a missed
323 // wakeup is impossible. Spurious unparks just loop.
324 if inner.epoch.load(Ordering::SeqCst) == seen
325 && !inner.shutdown.load(Ordering::Relaxed)
326 {
327 std::thread::park();
328 }
329 inner.parked[idx].store(false, Ordering::SeqCst);
330 }
331 }
332 // SAFETY: the slot was written before the epoch bump we just
333 // observed (release/acquire), and stays valid until `remaining`
334 // drops to zero — which happens only after `f` returns below.
335 let (task, n) = unsafe { (*inner.slot.get()).expect("job published with epoch") };
336 let f = unsafe { &*task.0 };
337 f(idx, n);
338 inner.remaining.fetch_sub(1, Ordering::AcqRel);
339 }
340}
341
342/// Row-parallel dense matvec: `out[o] = Σ_j w[o·in + j]·x[j]`.
343/// Bit-identical to the serial loop (row order does not change math).
344pub fn matvec_rows(pool: Option<&Pool>, w: &[f32], x: &[f32], out: &mut [f32]) {
345 let in_dim = x.len();
346 let out_dim = out.len();
347 debug_assert!(w.len() >= out_dim * in_dim);
348
349 let row_dot = |o: usize| -> f32 {
350 let row = &w[o * in_dim..(o + 1) * in_dim];
351 let mut sum = 0.0f32;
352 for j in 0..in_dim {
353 sum += row[j] * x[j];
354 }
355 sum
356 };
357
358 match pool {
359 // Small outputs are not worth the barrier round-trip.
360 Some(pool) if out_dim >= 256 => {
361 let out_addr = SendMut(out.as_mut_ptr());
362 pool.run(&move |widx, n| {
363 let chunk = out_dim.div_ceil(n);
364 let start = widx * chunk;
365 let end = (start + chunk).min(out_dim);
366 for o in start..end {
367 // SAFETY: workers write disjoint index ranges.
368 unsafe { *out_addr.at(o) = row_dot(o) };
369 }
370 });
371 }
372 _ => {
373 for (o, dst) in out.iter_mut().enumerate() {
374 *dst = row_dot(o);
375 }
376 }
377 }
378}
379
380/// Two-input row matvec: one pass over the weight rows serves BOTH
381/// inputs — CPU decode is memory-bound, so the second position costs a
382/// fraction of the first (this is where MTP speculative verify wins).
383/// Per-output accumulation order matches the single-input path exactly
384/// → bit-identical results.
385pub fn matvec_rows2(
386 pool: Option<&Pool>,
387 w: &[f32],
388 x1: &[f32],
389 x2: &[f32],
390 out1: &mut [f32],
391 out2: &mut [f32],
392) {
393 let in_dim = x1.len();
394 debug_assert_eq!(x2.len(), in_dim);
395 let out_dim = out1.len();
396 debug_assert_eq!(out2.len(), out_dim);
397 debug_assert!(w.len() >= out_dim * in_dim);
398
399 let row_dots = |o: usize| -> (f32, f32) {
400 let row = &w[o * in_dim..(o + 1) * in_dim];
401 let (mut s1, mut s2) = (0.0f32, 0.0f32);
402 for j in 0..in_dim {
403 s1 += row[j] * x1[j];
404 s2 += row[j] * x2[j];
405 }
406 (s1, s2)
407 };
408
409 match pool {
410 Some(pool) if out_dim >= 256 => {
411 let o1 = SendMut(out1.as_mut_ptr());
412 let o2 = SendMut(out2.as_mut_ptr());
413 pool.run(&move |widx, n| {
414 let chunk = out_dim.div_ceil(n);
415 let start = widx * chunk;
416 let end = (start + chunk).min(out_dim);
417 for o in start..end {
418 let (s1, s2) = row_dots(o);
419 // SAFETY: workers write disjoint index ranges.
420 unsafe {
421 *o1.at(o) = s1;
422 *o2.at(o) = s2;
423 }
424 }
425 });
426 }
427 _ => {
428 for o in 0..out_dim {
429 let (s1, s2) = row_dots(o);
430 out1[o] = s1;
431 out2[o] = s2;
432 }
433 }
434 }
435}
436
437#[derive(Clone, Copy)]
438struct SendMut(*mut f32);
439unsafe impl Send for SendMut {}
440unsafe impl Sync for SendMut {}
441
442impl SendMut {
443 /// Method receiver forces the closure to capture the whole (Sync)
444 /// wrapper, not the bare `*mut f32` field (edition-2021 precise capture).
445 #[inline]
446 fn at(self, i: usize) -> *mut f32 {
447 unsafe { self.0.add(i) }
448 }
449}
450
451#[cfg(test)]
452mod tests {
453 #[test]
454 fn capacity_split_clock_bins_vs_microarch() {
455 type P = super::Pool;
456 // JR510: all-A55, two clock bins — use every core.
457 assert_eq!(P::cores_from_capacities(&[1024, 1024, 1024, 1024, 768, 768, 768, 768]), Some(8));
458 // Classic big.LITTLE (A78 + A55) — big only.
459 assert_eq!(P::cores_from_capacities(&[1024, 1024, 1024, 1024, 350, 350, 350, 350]), Some(4));
460 // Three-tier flagship: X + A7xx mids stay, A5xx littles go.
461 assert_eq!(
462 P::cores_from_capacities(&[1024, 800, 800, 800, 800, 300, 300, 300]),
463 Some(5)
464 );
465 // Uniform: no signal, caller falls back.
466 assert_eq!(P::cores_from_capacities(&[1024; 8]), None);
467 assert_eq!(P::cores_from_capacities(&[]), None);
468 }
469
470 use super::*;
471
472 #[test]
473 fn parallel_matvec_equals_serial_bitexact() {
474 let (out_dim, in_dim) = (512, 64);
475 let w: Vec<f32> = (0..out_dim * in_dim).map(|i| (i as f32 * 0.013).sin()).collect();
476 let x: Vec<f32> = (0..in_dim).map(|i| (i as f32 * 0.07).cos()).collect();
477
478 let mut serial = vec![0.0f32; out_dim];
479 matvec_rows(None, &w, &x, &mut serial);
480
481 let pool = Pool::new(4);
482 let mut parallel = vec![0.0f32; out_dim];
483 matvec_rows(Some(&pool), &w, &x, &mut parallel);
484
485 assert_eq!(serial, parallel, "row-parallel must be bit-identical");
486 }
487
488 #[test]
489 fn fused_pair_equals_two_singles_bitexact() {
490 let (out_dim, in_dim) = (300, 48);
491 let w: Vec<f32> = (0..out_dim * in_dim).map(|i| (i as f32 * 0.011).sin()).collect();
492 let x1: Vec<f32> = (0..in_dim).map(|i| (i as f32 * 0.03).cos()).collect();
493 let x2: Vec<f32> = (0..in_dim).map(|i| (i as f32 * 0.09).sin()).collect();
494
495 let mut a1 = vec![0.0f32; out_dim];
496 let mut a2 = vec![0.0f32; out_dim];
497 matvec_rows(None, &w, &x1, &mut a1);
498 matvec_rows(None, &w, &x2, &mut a2);
499
500 for pool in [None, Some(Pool::new(3))] {
501 let mut b1 = vec![0.0f32; out_dim];
502 let mut b2 = vec![0.0f32; out_dim];
503 matvec_rows2(pool.as_ref(), &w, &x1, &x2, &mut b1, &mut b2);
504 assert_eq!(a1, b1, "fused lane 1 must be bit-identical");
505 assert_eq!(a2, b2, "fused lane 2 must be bit-identical");
506 }
507 }
508
509 #[test]
510 fn pool_survives_many_runs() {
511 let pool = Pool::new(3);
512 let counter = AtomicUsize::new(0);
513 for _ in 0..100 {
514 pool.run(&|_, _| {
515 counter.fetch_add(1, Ordering::Relaxed);
516 });
517 }
518 // 3 workers + the participating caller = 4 executions per run.
519 assert_eq!(counter.load(Ordering::Relaxed), 400);
520 }
521
522 #[test]
523 fn pool_wakes_after_park() {
524 // Force immediate parking (no spin) — the epoch/parked handshake
525 // must still never miss a wakeup.
526 let pool = Pool::with_spin(2, 0);
527 let counter = AtomicUsize::new(0);
528 for _ in 0..50 {
529 pool.run(&|_, _| {
530 counter.fetch_add(1, Ordering::Relaxed);
531 });
532 // Give workers time to actually park between jobs.
533 std::thread::sleep(std::time::Duration::from_micros(200));
534 }
535 assert_eq!(counter.load(Ordering::Relaxed), 150);
536 }
537
538 #[test]
539 fn worker_indices_are_distinct_and_cover_range() {
540 let pool = Pool::new(3);
541 let hits: Vec<AtomicUsize> = (0..4).map(|_| AtomicUsize::new(0)).collect();
542 for _ in 0..20 {
543 pool.run(&|widx, n| {
544 assert_eq!(n, 4);
545 hits[widx].fetch_add(1, Ordering::Relaxed);
546 });
547 }
548 for (i, h) in hits.iter().enumerate() {
549 assert_eq!(h.load(Ordering::Relaxed), 20, "participant {i} missed runs");
550 }
551 }
552}