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 one shared job descriptor + a per-worker ticket (roadmap
10//! §3 P0): the caller writes the descriptor, hands a ticket to each
11//! worker it invites and JOINS THE WORK as the extra worker instead of
12//! blocking on a latch. The previous design allocated an `Arc<Latch>`
13//! and pushed a message into every worker's mpsc channel for every
14//! matvec (~200 dispatches/token) — with decode-grade matvecs that
15//! synchronization was its own budget. Workers spin for
16//! `CMF_POOL_SPIN` iterations before parking.
17//! Default 4000: at ~39 dispatches/token, park-immediately pays the
18//! unpark syscall on every worker for every dispatch — measured on an
19//! M4 (interleaved A/B, current epoch dispatch + parked-flag design):
20//! Qwen-0.5B q8 decode 101→115 tok/s, q4t 117→149, the 50M bench model
21//! 549→954 at spin=4000 vs spin=0. An early measurement that showed
22//! spinning LOSING (−25% on q8) predates the parked-flag skip and the
23//! multi-matrix dispatch cuts; it no longer reproduces. Over-spinning
24//! still hurts (200k: −15% vs 4k — spinners steal the caller's serial
25//! cycles), so the budget stays bounded. `CMF_POOL_SPIN=0` restores
26//! park-immediately for share-the-box serving.
27//!
28//! `CMF_THREADS` env: 0/1 = serial, N = worker count
29//! (default: available_parallelism − 1, capped at 8).
30
31use std::sync::Arc;
32#[cfg(any(target_os = "android", target_os = "linux"))]
33use std::sync::Mutex;
34use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
35
36/// Embedder override for the pool size (C ABI `cortiq_set_threads`):
37/// 0 = unset, consult CMF_THREADS / topology as before. Read once at
38/// pool construction, so set it before the load.
39pub static FORCED_THREADS: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
40
41/// Kernel thread ids of the last fully constructed pool's workers
42/// (Android/Linux) — what ADPF's PerformanceHintManager needs to attribute
43/// work to the governor. Published as one complete snapshot after that
44/// pool's per-instance registration barrier; empty elsewhere.
45pub static WORKER_TIDS: std::sync::Mutex<Vec<i32>> = std::sync::Mutex::new(Vec::new());
46
47/// Keeps a word that one thread writes and others poll off everybody
48/// else's cache line (128 bytes: Apple silicon lines, and the adjacent-
49/// line prefetcher pair on x86).
50#[repr(align(128))]
51struct Padded<T>(T);
52
53/// One worker's mailbox, alone on its line: the worker polls it, only the
54/// caller writes `ticket`, only the worker writes `parked`.
55#[repr(align(128))]
56struct Slot {
57 /// Jobs handed to this worker so far. The caller bumps it (Release,
58 /// after the descriptor) once per job the worker is invited to; the
59 /// worker runs exactly one job per bump. The caller issues the next
60 /// bump only after this worker's `remaining` decrement for the
61 /// previous one, so the worker can never miss or double a ticket.
62 ticket: AtomicUsize,
63 /// "I am parked" — lets the caller skip the unpark syscall for a
64 /// worker that is still spinning.
65 parked: AtomicBool,
66}
67
68struct Inner {
69 /// Bumped once per published job, invited or not. Spinning workers
70 /// read it ONLY to re-arm their spin budget (the pool is busy, the
71 /// next job is microseconds away); it never makes a worker run
72 /// anything or read the descriptor.
73 epoch: AtomicUsize,
74 /// The published job: closure (fat-pointer halves), participant
75 /// count, publisher's GPU device. The device rides along because a
76 /// dispatch begun on card 1 must not finish on card 0: worker threads
77 /// have their own thread-locals, and the engine resolves its wgpu
78 /// context through one.
79 ///
80 /// Only workers holding a ticket for the current job read these
81 /// words, and the caller rewrites them only after `remaining` hit 0,
82 /// i.e. after every ticket holder has finished — so a read never
83 /// races a write and there is no torn descriptor to detect.
84 ///
85 /// History, because both earlier protocols failed: (1) one slot read
86 /// by EVERY worker on the epoch bump. A worker not invited to job k
87 /// is not waited for, so it could be preempted between seeing epoch k
88 /// and reading the slot, by which time the slot held job k+1 — it ran
89 /// k+1, decremented `remaining`, then saw epoch k+1 as new and ran it
90 /// AGAIN: `remaining` wrapped to `usize::MAX` and the caller spun
91 /// forever (the 47-minute `cortiq ppl` on the S4 bounded export,
92 /// whose 384-row Embryo matrices make every dispatch a limited one).
93 /// (2) The same words under a seqlock with the epoch inside: correct
94 /// on x86, but no release fence followed the writer's opening
95 /// increment, so the memory model did not order it before the word
96 /// stores (weakly ordered ARM may show a reader new words under an
97 /// old even count), and every limited dispatch still woke every
98 /// spinning worker to read the descriptor and skip it. With more
99 /// spinners than CPUs (the straggler test's 4× pool on a 2–4-CPU CI
100 /// runner) they held the CPUs while the one invited worker waited
101 /// for a time slice: ~10 ms per dispatch, 20 k jobs in ~230 s. A
102 /// worker now learns it is invited from its own ticket and nothing
103 /// else.
104 desc_data: AtomicUsize,
105 desc_vtable: AtomicUsize,
106 desc_n: AtomicUsize,
107 desc_dev: AtomicUsize,
108 /// Ticket holders still running the current job (excludes the
109 /// caller). Its own line: the caller polls it while workers
110 /// decrement.
111 remaining: Padded<AtomicUsize>,
112 /// One mailbox per worker, same order as `Pool::threads`.
113 slots: Box<[Slot]>,
114 shutdown: AtomicBool,
115 /// Spin iterations before a worker parks (0 = park immediately).
116 spin_budget: AtomicUsize,
117 /// More threads (workers + caller) than CPUs this process may run
118 /// on. Spinning cannot help then: a spinner occupies the CPU the
119 /// thread that has the work needs, and the job waits for the
120 /// scheduler's time slice. So in this mode a worker's spin budget is
121 /// re-armed only by its own tickets (an idle worker parks instead of
122 /// spinning through other workers' jobs), and every spinner — worker
123 /// or waiting caller — yields the CPU between short spin bursts.
124 /// The default size (available_parallelism − 1 workers) never is.
125 oversubscribed: bool,
126 /// Per-pool registration state. `WORKER_TIDS` is a process-wide
127 /// snapshot for ADPF and cannot be a construction barrier: another
128 /// pool may clear and republish that snapshot concurrently.
129 #[cfg(any(target_os = "android", target_os = "linux"))]
130 registered: AtomicUsize,
131 #[cfg(any(target_os = "android", target_os = "linux"))]
132 worker_tids: Mutex<Vec<i32>>,
133}
134
135/// Process-wide dispatch counter (roadmap §3 P0 «измерения»): one tick
136/// per published job. `bench --json` reports dispatches/token from it.
137static DISPATCHES: AtomicUsize = AtomicUsize::new(0);
138
139/// Total pool jobs published since process start (all pools).
140pub fn dispatch_count() -> usize {
141 DISPATCHES.load(Ordering::Relaxed)
142}
143
144/// Persistent thread pool: shared job slot, epoch dispatch, caller
145/// participation.
146pub struct Pool {
147 inner: Arc<Inner>,
148 /// Thread handles for `unpark` (same order as `parked`).
149 threads: Vec<std::thread::Thread>,
150 joins: Vec<std::thread::JoinHandle<()>>,
151}
152
153fn spin_budget_from_env() -> usize {
154 std::env::var("CMF_POOL_SPIN")
155 .ok()
156 .and_then(|v| v.parse::<usize>().ok())
157 .unwrap_or(4000)
158}
159
160/// Rows per chunk: enough chunks to balance, large enough to keep the SDOT
161/// inner loop and the prefetcher in their stride — and never so coarse that
162/// ONE worker takes the whole job.
163///
164/// That last clause was missing. The floor was a flat 32, so any job with
165/// fewer than 32 rows went entirely to whichever worker grabbed the cursor
166/// first while the other 48 were woken, found nothing, and left. The
167/// hyper-connection projection has 24 rows and is called 86 times a token:
168/// it paid the full price of a fan-out and ran single-threaded.
169pub(crate) fn grain_for(rows: usize, workers: usize) -> usize {
170 if rows == 0 || workers <= 1 {
171 return rows.max(1);
172 }
173 let balanced = (rows / (workers * 8)).max(32);
174 // One chunk per worker at the very least.
175 balanced.min(rows.div_ceil(workers)).max(1)
176}
177
178impl Pool {
179 pub fn new(n_workers: usize) -> Self {
180 Self::with_spin(n_workers, spin_budget_from_env())
181 }
182
183 /// Explicit spin budget (tests pin it without touching the env).
184 pub fn with_spin(n_workers: usize, spin_budget: usize) -> Self {
185 // Affinity- and cgroup-quota-aware on Linux; read once — the NUMA
186 // bind never narrows the mask below the pool's thread count.
187 let cpus = std::thread::available_parallelism()
188 .map(|n| n.get())
189 .unwrap_or(1);
190 let inner = Arc::new(Inner {
191 epoch: AtomicUsize::new(0),
192 desc_data: AtomicUsize::new(0),
193 desc_vtable: AtomicUsize::new(0),
194 desc_n: AtomicUsize::new(0),
195 desc_dev: AtomicUsize::new(0),
196 remaining: Padded(AtomicUsize::new(0)),
197 slots: (0..n_workers)
198 .map(|_| Slot {
199 ticket: AtomicUsize::new(0),
200 parked: AtomicBool::new(false),
201 })
202 .collect(),
203 shutdown: AtomicBool::new(false),
204 spin_budget: AtomicUsize::new(spin_budget),
205 oversubscribed: n_workers + 1 > cpus,
206 #[cfg(any(target_os = "android", target_os = "linux"))]
207 registered: AtomicUsize::new(0),
208 #[cfg(any(target_os = "android", target_os = "linux"))]
209 worker_tids: Mutex::new(Vec::with_capacity(n_workers)),
210 });
211 let mut joins = Vec::with_capacity(n_workers);
212 for w in 0..n_workers {
213 let inner = inner.clone();
214 let h = std::thread::Builder::new()
215 .name(format!("cmf-pool-{w}"))
216 .spawn(move || {
217 #[cfg(any(target_os = "android", target_os = "linux"))]
218 {
219 let tid = unsafe { libc::gettid() } as i32;
220 if let Ok(mut tids) = inner.worker_tids.lock() {
221 tids.push(tid);
222 }
223 inner.registered.fetch_add(1, Ordering::Release);
224 }
225 worker_loop(&inner, w)
226 })
227 .expect("spawn pool worker");
228 joins.push(h);
229 }
230 // Per-pool registration barrier: `spawn` returns before the closure runs,
231 // and the embedder reads `cortiq_worker_tids` right after load —
232 // on a phone only the first worker had registered by then (the
233 // '· 1 threads' About line that misled the cmfmobile device
234 // investigation twice). Thread start is milliseconds; wait for
235 // every worker has registered before construction returns.
236 #[cfg(any(target_os = "android", target_os = "linux"))]
237 while inner.registered.load(Ordering::Acquire) < n_workers {
238 std::thread::yield_now();
239 }
240 #[cfg(any(target_os = "android", target_os = "linux"))]
241 if let (Ok(mut global), Ok(local)) = (WORKER_TIDS.lock(), inner.worker_tids.lock()) {
242 *global = local.clone();
243 }
244 let threads = joins.iter().map(|h| h.thread().clone()).collect();
245 Self {
246 inner,
247 threads,
248 joins,
249 }
250 }
251
252 /// Big-core count on heterogeneous ARM (big.LITTLE): the kernel
253 /// exposes per-core capacity on Android and most ARM Linux; efficiency
254 /// cores in the pool DRAG the big ones on our row-parallel jobs (the
255 /// same cliff llama.cpp hits at -t 10 on an M4: 163 → 112 tok/s).
256 /// None = capacities absent or homogeneous.
257 #[cfg(all(
258 target_arch = "aarch64",
259 any(target_os = "linux", target_os = "android")
260 ))]
261 fn big_cores() -> Option<usize> {
262 Self::cores_from_capacities(&core_capacities())
263 }
264
265 /// How many cores the pool should use, from the kernel's per-core
266 /// capacity values. Capacity folds µarch × clock into one number,
267 /// and the two need different treatment: cores of ANOTHER µarch
268 /// (A5xx efficiency cluster next to A7xx/X: capacity ratio ≥ ~2)
269 /// drag row-parallel work down and are excluded; cores of the SAME
270 /// µarch merely clock-binned (JLQ JR510: 8×A55 as 4×2.0 + 4×1.5 GHz,
271 /// ratio 1.33) pull their weight and must ALL be used. The 1.6
272 /// threshold splits the two regimes: on a Snapdragon 8-class part
273 /// it keeps X + A7xx mid cores and drops A5xx.
274 #[cfg_attr(
275 not(all(
276 target_arch = "aarch64",
277 any(target_os = "linux", target_os = "android")
278 )),
279 allow(dead_code)
280 )]
281 fn cores_from_capacities(caps: &[u64]) -> Option<usize> {
282 let max = *caps.iter().max()?;
283 let min = *caps.iter().min()?;
284 if caps.len() < 2 || max == min {
285 return None;
286 }
287 Some(caps.iter().filter(|&&c| c * 8 >= max * 5).count())
288 }
289
290 #[cfg(target_os = "macos")]
291 fn big_cores() -> Option<usize> {
292 // Apple silicon: the P-only default measured WORSE than mixing the
293 // efficiency cores in — the grain-pulling dispatch absorbs the
294 // speed skew exactly as designed, and decode is memory-bound
295 // enough that E-cores add real serviceable work (M4, dense 3B:
296 // 4 threads 8.4 tok/s, 6-9 threads 9.6-10.7). Fall through to
297 // available_parallelism - 1; CMF_THREADS still pins by hand.
298 // The sysctl probe stays for introspection tooling.
299 if true {
300 return None;
301 }
302 #[allow(unreachable_code)]
303 unsafe extern "C" {
304 fn sysctlbyname(
305 name: *const std::ffi::c_char,
306 oldp: *mut std::ffi::c_void,
307 oldlenp: *mut usize,
308 newp: *mut std::ffi::c_void,
309 newlen: usize,
310 ) -> std::ffi::c_int;
311 }
312 unsafe {
313 let name = std::ffi::CString::new("hw.perflevel0.physicalcpu").ok()?;
314 let mut count: i32 = 0;
315 let mut size = std::mem::size_of::<i32>();
316 let ret = sysctlbyname(
317 name.as_ptr(),
318 &mut count as *mut i32 as *mut std::ffi::c_void,
319 &mut size,
320 std::ptr::null_mut(),
321 0,
322 );
323 if ret == 0 && count > 0 {
324 Some(count as usize)
325 } else {
326 None
327 }
328 }
329 }
330
331 #[cfg(not(any(
332 all(
333 target_arch = "aarch64",
334 any(target_os = "linux", target_os = "android")
335 ),
336 target_os = "macos"
337 )))]
338 fn big_cores() -> Option<usize> {
339 None
340 }
341
342 /// The thread count `from_env` would use RIGHT NOW: forced (C ABI)
343 /// > CMF_THREADS > big-core topology > available_parallelism−1.
344 /// > ≤1 means the model runs serial (no pool). Introspection
345 /// > (`execution_mode`, status endpoints) must report THIS, not
346 /// > available_parallelism.
347 pub fn effective_threads() -> usize {
348 let forced = FORCED_THREADS.load(std::sync::atomic::Ordering::Relaxed);
349 if forced > 0 {
350 return forced;
351 }
352 match std::env::var("CMF_THREADS") {
353 Ok(v) => v.parse::<usize>().unwrap_or(0),
354 Err(_) => match Self::big_cores() {
355 Some(big) => big,
356 None => {
357 // The cap was 8, which left big machines idle: on a
358 // 256-core EPYC, Nanbeige 4.2 decoded at 7.4 tok/s on
359 // the default 8 threads and 14.8 at 32, with prefill
360 // 12 -> ~16 over the same move. Past ~32 it falls off
361 // hard (5.5 at 64, 1.6 at 256) — decode is
362 // memory-bound and the extra threads only add
363 // dispatch barriers — so 32 is a ceiling, not a
364 // target. Machines with 9 cores or fewer are
365 // unaffected: avail-1 already bounds them.
366 let avail = std::thread::available_parallelism()
367 .map(|n| n.get())
368 .unwrap_or(1);
369 avail.saturating_sub(1).min(32)
370 }
371 },
372 }
373 }
374
375 /// Pool sized from `CMF_THREADS` (see module docs). `None` = serial.
376 /// Without the env, heterogeneous ARM defaults to its BIG cores.
377 pub fn from_env() -> Option<Arc<Self>> {
378 let n = Self::effective_threads();
379 if n <= 1 {
380 None
381 } else {
382 Some(Arc::new(Self::new(n)))
383 }
384 }
385
386 /// Spawned worker threads (the caller joins each job on top).
387 pub fn n_workers(&self) -> usize {
388 self.threads.len()
389 }
390
391 /// Keep the pool on the NUMA node that holds `regions` (the model's
392 /// weight bytes). Linux with two or more nodes only; `CMF_NUMA=0`
393 /// turns it off, `CMF_NUMA=node:<n>` forces a node.
394 ///
395 /// WHY: decode streams every weight once per token, and on a
396 /// two-socket host the page cache holds a file on whichever node
397 /// read it. Unpinned, the scheduler spreads the workers over both
398 /// sockets and half the matvec rows cross the socket link. Measured
399 /// on a 2×EPYC 7763 pod with the model's pages all on node 0 (31 CPUs
400 /// of cgroup quota): a STREAM-style read over a node-0 buffer gives
401 /// 42 GB/s from 31 unpinned threads and 74 GB/s from 31 threads kept
402 /// on node 0. The mask is the node's physical cores (first SMT
403 /// sibling) when there are enough of them for the pool, else the
404 /// whole node; never narrower than the pool, so nothing oversubscribes.
405 /// Threads are bound to a SET of cores, not to one core each: the
406 /// scheduler still balances inside the node. The calling thread
407 /// adopts the same mask on its next dispatch.
408 pub fn bind_numa(&self, regions: &[&[u8]]) {
409 #[cfg(target_os = "linux")]
410 {
411 let Some((node, cpus)) = numa::choose(regions, self.threads.len() + 1) else {
412 return;
413 };
414 let mut applied = 0usize;
415 if let Ok(tids) = self.inner.worker_tids.lock() {
416 for &tid in tids.iter() {
417 if numa::set_affinity(tid, &cpus) {
418 applied += 1;
419 }
420 }
421 }
422 numa::publish(cpus.clone());
423 numa::adopt_caller();
424 tracing::info!(
425 "numa: pool bound to node {node} ({} cpus, {applied}/{} workers)",
426 cpus.len(),
427 self.threads.len()
428 );
429 if std::env::var("CMF_NUMA_TRACE").is_ok_and(|v| v != "0") {
430 eprintln!(
431 "numa: pool bound to node {node}: {} cpus, {applied}/{} workers",
432 cpus.len(),
433 self.threads.len()
434 );
435 }
436 }
437 #[cfg(not(target_os = "linux"))]
438 let _ = regions;
439 }
440
441 /// Retune an already-created pool for an architecture with a measured
442 /// dispatch cadence. The environment remains the operator override; this
443 /// hook only changes the automatic default after model geometry is known.
444 pub(crate) fn set_spin_budget(&self, spins: usize) {
445 self.inner.spin_budget.store(spins, Ordering::Relaxed);
446 }
447
448 /// One job: write the descriptor, hand a ticket to workers `0..nw`,
449 /// run the caller's share as participant `nw` of `nw + 1`, and return
450 /// once every ticket holder has finished. Only one job is ever in
451 /// flight (this drains `remaining` before it returns), so the caller
452 /// is the single writer of the descriptor, `remaining` and every
453 /// `ticket`; workers without a ticket never touch any of them.
454 fn dispatch(&self, f: &(dyn Fn(usize, usize) + Sync), nw: usize) {
455 let inner = &*self.inner;
456 let n = nw + 1;
457 let ptr: *const (dyn Fn(usize, usize) + Sync) = f;
458 // SAFETY: a fat pointer is exactly two words on every supported
459 // target; the halves are only ever reassembled by `worker_loop`.
460 let raw: [usize; 2] = unsafe { std::mem::transmute(ptr) };
461 inner.desc_data.store(raw[0], Ordering::Relaxed);
462 inner.desc_vtable.store(raw[1], Ordering::Relaxed);
463 inner.desc_n.store(n, Ordering::Relaxed);
464 inner
465 .desc_dev
466 .store(crate::gpu::current_device(), Ordering::Relaxed);
467 inner.remaining.0.store(nw, Ordering::Relaxed);
468 let e = inner.epoch.load(Ordering::Relaxed);
469 inner.epoch.store(e.wrapping_add(1), Ordering::Relaxed);
470 // Release: a worker that sees its new ticket (Acquire) sees the
471 // descriptor and `remaining` written above.
472 let slots = &inner.slots[..nw];
473 for slot in slots {
474 let t = slot.ticket.load(Ordering::Relaxed).wrapping_add(1);
475 slot.ticket.store(t, Ordering::Release);
476 }
477 // Pairs with the fence a worker issues between raising `parked`
478 // and re-reading its ticket: either that worker sees the ticket
479 // or we see its flag and unpark it — no lost wakeup. One fence
480 // for all the tickets instead of a SeqCst store per worker.
481 std::sync::atomic::fence(Ordering::SeqCst);
482 for (slot, t) in slots.iter().zip(&self.threads) {
483 if slot.parked.load(Ordering::Relaxed) {
484 t.unpark();
485 }
486 }
487
488 // The caller's share — the barrier costs nothing while there is
489 // real work to do.
490 f(nw, n);
491
492 // Wait for the stragglers (bounded by one worker's chunk). When
493 // the pool has more threads than CPUs, a ticket holder may be
494 // waiting for THIS core: yield it early instead of spinning.
495 let spin_limit = if inner.oversubscribed { 64 } else { 10_000 };
496 let mut spins = 0usize;
497 while inner.remaining.0.load(Ordering::Acquire) != 0 {
498 spins += 1;
499 if spins < spin_limit {
500 std::hint::spin_loop();
501 } else {
502 std::thread::yield_now();
503 }
504 }
505 }
506
507 /// Run `f(row_start, row_end)` over `0..rows`, self-balancing.
508 ///
509 /// One dispatch, but workers pull row-ranges from a shared cursor
510 /// instead of each taking a fixed 1/n slice. On a heterogeneous CPU
511 /// (Apple Silicon: 4 P-cores + 6 E-cores here) a static split makes
512 /// every matvec end at the SLOWEST core's pace while the fast ones
513 /// idle at the barrier; pulling by grain lets a P-core take several
514 /// chunks for each one an E-core takes, so skew collapses to a
515 /// single grain. Row ranges stay disjoint and each row's dot is
516 /// computed exactly as in the serial path → bit-identical output.
517 pub fn run_rows(&self, rows: usize, f: &(dyn Fn(usize, usize) + Sync)) {
518 let grain = grain_for(rows, self.threads.len() + 1);
519 let chunks = rows.div_ceil(grain.max(1));
520 let next = AtomicUsize::new(0);
521 self.run_limited(chunks, &|_w, _n| loop {
522 let start = next.fetch_add(grain, Ordering::Relaxed);
523 if start >= rows {
524 break;
525 }
526 f(start, (start + grain).min(rows));
527 });
528 }
529
530 /// `run`, with at most `max_workers` workers PARTICIPATING. Same
531 /// grain, same row split, bit-identical results — only the number of
532 /// threads handed the job changes: a job with eight grains has no use
533 /// for three hundred workers, the unpark syscalls and the
534 /// remaining-drain would BE the job (measured: 361 pool dispatches
535 /// per DeepSeek-V4 token, and CMF_THREADS=64 vs 380 was 1.3 vs 2.4
536 /// tok/s with no other change). Only cursor-style closures (which
537 /// ignore their (idx, n) arguments) come through here: the caller
538 /// identifies itself as the capped count, which is NOT `n_workers()`.
539 fn run_limited(&self, max_workers: usize, f: &(dyn Fn(usize, usize) + Sync)) {
540 let nw = self.threads.len().min(max_workers);
541 if nw == self.threads.len() {
542 return self.run(f);
543 }
544 #[cfg(target_os = "linux")]
545 numa::adopt_caller();
546 DISPATCHES.fetch_add(1, Ordering::Relaxed);
547 // Workers `nw..` get no ticket: they neither run nor wait.
548 self.dispatch(f, nw);
549 }
550
551 /// Multi-matrix job: one dispatch serves SEVERAL row spaces
552 /// (roadmap §3 P0 — «одна внешняя публикация job на слой»). Parts
553 /// are laid out back-to-back in a virtual row space and pulled by
554 /// grain from one shared cursor, so QKV or gate+up cost a single
555 /// barrier instead of one each. Each part's `f(start, end)` sees its
556 /// OWN row indices — per-row math and outputs are bit-identical to
557 /// separate `run_rows` calls.
558 pub fn run_many(&self, parts: &[(usize, &(dyn Fn(usize, usize) + Sync))]) {
559 let total: usize = parts.iter().map(|p| p.0).sum();
560 if total == 0 {
561 return;
562 }
563 let grain = grain_for(total, self.threads.len() + 1);
564 let chunks = total.div_ceil(grain.max(1));
565 let next = AtomicUsize::new(0);
566 self.run_limited(chunks, &|_w, _n| loop {
567 let s = next.fetch_add(grain, Ordering::Relaxed);
568 if s >= total {
569 break;
570 }
571 let e = (s + grain).min(total);
572 let mut base = 0usize;
573 for &(rows, f) in parts {
574 let a = s.max(base);
575 let b = e.min(base + rows);
576 if a < b {
577 f(a - base, b - base);
578 }
579 base += rows;
580 if base >= e {
581 break;
582 }
583 }
584 });
585 }
586
587 /// Run `f(worker_idx, n_participants)` on every worker AND the
588 /// calling thread (`worker_idx = n_workers()` for the caller);
589 /// returns when all participants have finished.
590 pub fn run(&self, f: &(dyn Fn(usize, usize) + Sync)) {
591 #[cfg(target_os = "linux")]
592 numa::adopt_caller();
593 DISPATCHES.fetch_add(1, Ordering::Relaxed);
594 self.dispatch(f, self.threads.len());
595 }
596}
597
598impl Drop for Pool {
599 fn drop(&mut self) {
600 self.inner.shutdown.store(true, Ordering::SeqCst);
601 for t in &self.threads {
602 t.unpark();
603 }
604 for h in self.joins.drain(..) {
605 let _ = h.join();
606 }
607 }
608}
609
610/// NUMA placement for the pool (see `Pool::bind_numa`).
611#[cfg(target_os = "linux")]
612mod numa {
613 use std::sync::Mutex;
614 use std::sync::atomic::{AtomicUsize, Ordering};
615
616 /// The published mask; `EPOCH` bumps on every publish so a calling
617 /// thread re-adopts at most once per bind.
618 static MASK: Mutex<Vec<usize>> = Mutex::new(Vec::new());
619 static EPOCH: AtomicUsize = AtomicUsize::new(0);
620 thread_local! {
621 static SEEN: std::cell::Cell<usize> = const { std::cell::Cell::new(0) };
622 }
623
624 pub(super) fn publish(cpus: Vec<usize>) {
625 if let Ok(mut m) = MASK.lock() {
626 *m = cpus;
627 }
628 EPOCH.fetch_add(1, Ordering::Release);
629 }
630
631 /// One relaxed load + one TLS read per dispatch when nothing changed.
632 #[inline]
633 pub(super) fn adopt_caller() {
634 let e = EPOCH.load(Ordering::Acquire);
635 if e == 0 || SEEN.with(|c| c.get()) == e {
636 return;
637 }
638 SEEN.with(|c| c.set(e));
639 if let Ok(m) = MASK.lock() {
640 if !m.is_empty() {
641 set_affinity(0, &m);
642 }
643 }
644 }
645
646 pub(super) fn parse_list(s: &str) -> Vec<usize> {
647 let mut out = Vec::new();
648 for part in s.trim().split(',') {
649 let part = part.trim();
650 if part.is_empty() {
651 continue;
652 }
653 match part.split_once('-') {
654 Some((a, b)) => {
655 if let (Ok(a), Ok(b)) = (a.parse::<usize>(), b.parse::<usize>()) {
656 out.extend(a..=b);
657 }
658 }
659 None => {
660 if let Ok(a) = part.parse() {
661 out.push(a);
662 }
663 }
664 }
665 }
666 out
667 }
668
669 fn nodes() -> Vec<(usize, Vec<usize>)> {
670 let mut v = Vec::new();
671 let Ok(rd) = std::fs::read_dir("/sys/devices/system/node") else {
672 return v;
673 };
674 for e in rd.flatten() {
675 let name = e.file_name().to_string_lossy().to_string();
676 let Some(id) = name.strip_prefix("node").and_then(|x| x.parse::<usize>().ok()) else {
677 continue;
678 };
679 if let Ok(l) = std::fs::read_to_string(e.path().join("cpulist")) {
680 let cpus = parse_list(&l);
681 if !cpus.is_empty() {
682 v.push((id, cpus));
683 }
684 }
685 }
686 v.sort();
687 v
688 }
689
690 fn allowed() -> Vec<usize> {
691 // SAFETY: plain syscall into a zeroed, correctly sized set.
692 unsafe {
693 let mut set: libc::cpu_set_t = std::mem::zeroed();
694 if libc::sched_getaffinity(0, std::mem::size_of::<libc::cpu_set_t>(), &mut set) != 0 {
695 return Vec::new();
696 }
697 (0..libc::CPU_SETSIZE as usize)
698 .filter(|&c| libc::CPU_ISSET(c, &set))
699 .collect()
700 }
701 }
702
703 /// First SMT sibling of its core (or no topology info: count it).
704 fn primary(cpu: usize) -> bool {
705 let p = format!("/sys/devices/system/cpu/cpu{cpu}/topology/thread_siblings_list");
706 match std::fs::read_to_string(p) {
707 Ok(l) => parse_list(&l).first().is_none_or(|&f| f == cpu),
708 Err(_) => true,
709 }
710 }
711
712 /// Where the weights live: (pages sampled, sampled pages in the page
713 /// cache, mapped pages per node). The sample is ≤ 4096 pages spread
714 /// over `regions`; every sampled page that is already cached is mapped
715 /// here with one read (a minor fault — `mincore` says it is cached, so
716 /// no disk I/O), because both node queries below only see pages mapped
717 /// into THIS process. Per-node counts come from `move_pages` in query
718 /// mode, or — where a container's seccomp profile refuses that syscall
719 /// (EPERM on the RunPod image) — from `/proc/self/numa_maps` for the
720 /// mappings that hold the regions.
721 fn page_nodes(regions: &[&[u8]]) -> (usize, usize, Vec<usize>) {
722 const PAGE: usize = 4096;
723 let total: usize = regions.iter().map(|r| r.len() / PAGE).sum();
724 if total == 0 {
725 return (0, 0, Vec::new());
726 }
727 let stride = total.div_ceil(4096).max(1);
728 let mut pages: Vec<*mut libc::c_void> = Vec::new();
729 for r in regions {
730 let base = (r.as_ptr() as usize).div_ceil(PAGE) * PAGE;
731 let end = r.as_ptr() as usize + r.len();
732 let mut a = base;
733 while a + PAGE <= end {
734 pages.push(a as *mut libc::c_void);
735 a += PAGE * stride;
736 }
737 }
738 let mut incore = 0usize;
739 for &p in &pages {
740 let mut vec = 0u8;
741 // SAFETY: `p` is a page-aligned address inside a live mapping.
742 let cached = unsafe { libc::mincore(p, PAGE, &mut vec) } == 0 && vec & 1 == 1;
743 if cached {
744 incore += 1;
745 // SAFETY: readable mapped byte; volatile so it is not elided.
746 unsafe { std::ptr::read_volatile(p as *const u8) };
747 }
748 }
749 let mut status = vec![-1i32; pages.len()];
750 // SAFETY: query-only move_pages on our own mappings; `nodes` is
751 // NULL so nothing moves, `status` has one slot per page.
752 let rc = unsafe {
753 libc::syscall(
754 libc::SYS_move_pages,
755 0,
756 pages.len() as libc::c_ulong,
757 pages.as_mut_ptr(),
758 std::ptr::null::<libc::c_int>(),
759 status.as_mut_ptr(),
760 0,
761 )
762 };
763 let mut by = Vec::new();
764 if rc == 0 {
765 for &st in &status {
766 if st >= 0 {
767 let n = st as usize;
768 if by.len() <= n {
769 by.resize(n + 1, 0);
770 }
771 by[n] += 1;
772 }
773 }
774 } else {
775 by = numa_maps_nodes(regions);
776 }
777 (pages.len(), incore, by)
778 }
779
780 /// Mapped pages per node of every mapping that overlaps `regions`,
781 /// from `/proc/self/maps` (ranges) + `/proc/self/numa_maps` (`N<k>=`).
782 fn numa_maps_nodes(regions: &[&[u8]]) -> Vec<usize> {
783 let (Ok(maps), Ok(nm)) = (
784 std::fs::read_to_string("/proc/self/maps"),
785 std::fs::read_to_string("/proc/self/numa_maps"),
786 ) else {
787 return Vec::new();
788 };
789 let spans: Vec<(usize, usize)> = regions
790 .iter()
791 .map(|r| (r.as_ptr() as usize, r.as_ptr() as usize + r.len()))
792 .collect();
793 let mut starts = std::collections::HashSet::new();
794 for line in maps.lines() {
795 let Some((range, _)) = line.split_once(' ') else {
796 continue;
797 };
798 let Some((a, b)) = range.split_once('-') else {
799 continue;
800 };
801 let (Ok(a), Ok(b)) = (usize::from_str_radix(a, 16), usize::from_str_radix(b, 16)) else {
802 continue;
803 };
804 if spans.iter().any(|&(s, e)| s < b && a < e) {
805 starts.insert(a);
806 }
807 }
808 let mut by = Vec::new();
809 for line in nm.lines() {
810 let mut it = line.split_whitespace();
811 let Some(a) = it.next().and_then(|a| usize::from_str_radix(a, 16).ok()) else {
812 continue;
813 };
814 if !starts.contains(&a) {
815 continue;
816 }
817 for f in it {
818 let Some((k, v)) = f.split_once('=') else {
819 continue;
820 };
821 let (Some(n), Ok(v)) = (
822 k.strip_prefix('N').and_then(|n| n.parse::<usize>().ok()),
823 v.parse::<usize>(),
824 ) else {
825 continue;
826 };
827 if by.len() <= n {
828 by.resize(n + 1, 0);
829 }
830 by[n] += v;
831 }
832 }
833 by
834 }
835
836 /// (node, cpu mask) for a pool of `threads` participants, or None.
837 pub(super) fn choose(regions: &[&[u8]], threads: usize) -> Option<(usize, Vec<usize>)> {
838 let env = std::env::var("CMF_NUMA").ok();
839 if matches!(env.as_deref(), Some("0") | Some("off")) {
840 return None;
841 }
842 let trace = std::env::var("CMF_NUMA_TRACE").is_ok_and(|v| v != "0");
843 let nodes = nodes();
844 if nodes.len() < 2 {
845 if trace {
846 eprintln!("numa: {} node(s) visible — nothing to bind", nodes.len());
847 }
848 return None;
849 }
850 // `CMF_NUMA=node:<n>` forces a node (plain "0" means OFF).
851 let forced = env
852 .as_deref()
853 .and_then(|v| v.strip_prefix("node:"))
854 .and_then(|v| v.parse::<usize>().ok());
855 let node = match forced {
856 Some(n) => n,
857 None => {
858 // Auto: only when the weights already sit on ONE node
859 // (≥ 90% of the resident sample, and most of the sample
860 // resident). A file spread over both nodes is better
861 // served by both sockets; a cold file has no home yet.
862 let (sampled, incore, by) = page_nodes(regions);
863 let resident: usize = by.iter().sum();
864 if trace {
865 eprintln!(
866 "numa: sampled {sampled} weight pages, {incore} cached, mapped by node {by:?}"
867 );
868 }
869 if sampled == 0 || incore * 2 < sampled || resident == 0 {
870 return None;
871 }
872 let (n, &cnt) = by.iter().enumerate().max_by_key(|(_, c)| **c)?;
873 if cnt * 10 < resident * 9 {
874 return None;
875 }
876 n
877 }
878 };
879 let cpus = &nodes.iter().find(|(id, _)| *id == node)?.1;
880 let allowed = allowed();
881 let usable: Vec<usize> = cpus.iter().copied().filter(|c| allowed.contains(c)).collect();
882 let prim: Vec<usize> = usable.iter().copied().filter(|&c| primary(c)).collect();
883 if prim.len() >= threads {
884 Some((node, prim))
885 } else if usable.len() >= threads {
886 Some((node, usable))
887 } else {
888 None
889 }
890 }
891
892 /// Bind thread `tid` (0 = the calling thread) to `cpus`.
893 pub(super) fn set_affinity(tid: i32, cpus: &[usize]) -> bool {
894 // SAFETY: plain syscall with a zeroed, correctly sized set.
895 unsafe {
896 let mut set: libc::cpu_set_t = std::mem::zeroed();
897 for &c in cpus {
898 if c < libc::CPU_SETSIZE as usize {
899 libc::CPU_SET(c, &mut set);
900 }
901 }
902 libc::sched_setaffinity(tid, std::mem::size_of::<libc::cpu_set_t>(), &set) == 0
903 }
904 }
905}
906
907/// Per-core capacity: the kernel's `cpu_capacity` (µarch × clock) when
908/// EAS exposes it, else `cpufreq/cpuinfo_max_freq` — same cluster
909/// ordering, so the 62.5% big-core rule keeps working on EAS-less
910/// kernels (TUNING.md open item: pinning silently did nothing there).
911#[cfg(any(
912 target_os = "android",
913 all(target_arch = "aarch64", target_os = "linux")
914))]
915fn core_capacities() -> Vec<u64> {
916 let read_all = |leaf: &str| -> Vec<u64> {
917 let mut vals = Vec::new();
918 for cpu in 0.. {
919 let path = format!("/sys/devices/system/cpu/cpu{cpu}/{leaf}");
920 match std::fs::read_to_string(&path) {
921 Ok(v) => match v.trim().parse() {
922 Ok(x) => vals.push(x),
923 Err(_) => break,
924 },
925 Err(_) => break,
926 }
927 }
928 vals
929 };
930 let caps = read_all("cpu_capacity");
931 if caps.len() >= 2 {
932 return caps;
933 }
934 read_all("cpufreq/cpuinfo_max_freq")
935}
936
937#[cfg(target_os = "android")]
938fn pin_thread_to_big_cores() {
939 use std::mem;
940 let caps = core_capacities();
941 let max = caps.iter().copied().max().unwrap_or(0);
942 let min = caps.iter().copied().min().unwrap_or(0);
943
944 // Only pin if heterogeneous
945 if caps.len() < 2 || max == min {
946 return;
947 }
948
949 unsafe {
950 let mut set: libc::cpu_set_t = mem::zeroed();
951 for (i, &c) in caps.iter().enumerate() {
952 if c * 8 >= max * 5 {
953 libc::CPU_SET(i, &mut set);
954 }
955 }
956 libc::sched_setaffinity(0, mem::size_of::<libc::cpu_set_t>(), &set);
957 }
958}
959
960fn worker_loop(inner: &Inner, idx: usize) {
961 #[cfg(target_os = "android")]
962 pin_thread_to_big_cores();
963 // Apple silicon: ask for the performance cores. Threads spawned
964 // without a QoS class land on the efficiency cores when the
965 // scheduler feels like it — a user's video-VAE encode on an M4 sat
966 // on the E-cores at 100% with the P-cores asleep for 140 s (HF
967 // discussion #4). USER_INITIATED is the class an interactive tool's
968 // work belongs to; the ~4 P-cores then take the pool's grains.
969 #[cfg(target_os = "macos")]
970 unsafe {
971 libc::pthread_set_qos_class_self_np(libc::qos_class_t::QOS_CLASS_USER_INITIATED, 0);
972 }
973
974 let slot = &inner.slots[idx];
975 // Baselines are the construction values (0), never a fresh read: if
976 // the caller hands out a ticket before the OS actually starts this
977 // thread, adopting the live value as "already seen" would skip that
978 // job and deadlock the caller's wait.
979 let mut seen = 0usize;
980 let mut seen_epoch = 0usize;
981 loop {
982 // Wait for a ticket: spin first (decode publishes the next matvec
983 // within microseconds), park only when idle for real.
984 let mut spins = 0usize;
985 loop {
986 let t = slot.ticket.load(Ordering::Acquire);
987 if t != seen {
988 seen = t;
989 break;
990 }
991 if inner.shutdown.load(Ordering::Relaxed) {
992 return;
993 }
994 if !inner.oversubscribed {
995 // Any job — even one this worker sits out — means the
996 // pool is busy: stay hot for the next one.
997 let e = inner.epoch.load(Ordering::Relaxed);
998 if e != seen_epoch {
999 seen_epoch = e;
1000 spins = 0;
1001 }
1002 }
1003 if spins < inner.spin_budget.load(Ordering::Relaxed) {
1004 spins += 1;
1005 if inner.oversubscribed && spins.is_multiple_of(64) {
1006 std::thread::yield_now();
1007 } else {
1008 std::hint::spin_loop();
1009 }
1010 } else {
1011 slot.parked.store(true, Ordering::Relaxed);
1012 // Pairs with the caller's fence between writing tickets
1013 // and reading `parked`: either it sees our flag (and
1014 // unparks) or we see its ticket here — a missed wakeup is
1015 // impossible. Spurious unparks just loop.
1016 std::sync::atomic::fence(Ordering::SeqCst);
1017 if slot.ticket.load(Ordering::Relaxed) == seen
1018 && !inner.shutdown.load(Ordering::Relaxed)
1019 {
1020 std::thread::park();
1021 }
1022 slot.parked.store(false, Ordering::Relaxed);
1023 }
1024 }
1025 // The ticket's Acquire made the descriptor visible, and the caller
1026 // cannot rewrite it before our decrement below: it waits for
1027 // `remaining`, which counts this ticket.
1028 let data = inner.desc_data.load(Ordering::Relaxed);
1029 let vtable = inner.desc_vtable.load(Ordering::Relaxed);
1030 let n = inner.desc_n.load(Ordering::Relaxed);
1031 let dev = inner.desc_dev.load(Ordering::Relaxed);
1032 // SAFETY: the two words are the fat pointer `dispatch` split, and
1033 // the caller of this job is blocked on our decrement below, so the
1034 // closure it borrows is alive for the whole call.
1035 let task: *const (dyn Fn(usize, usize) + Sync + 'static) =
1036 unsafe { std::mem::transmute([data, vtable]) };
1037 let f = unsafe { &*task };
1038 crate::gpu::set_current_device(dev);
1039 f(idx, n);
1040 inner.remaining.0.fetch_sub(1, Ordering::AcqRel);
1041 }
1042}
1043
1044/// Row-parallel dense matvec: `out[o] = Σ_j w[o·in + j]·x[j]`.
1045/// Bit-identical to the serial loop (row order does not change math).
1046pub fn matvec_rows(pool: Option<&Pool>, w: &[f32], x: &[f32], out: &mut [f32]) {
1047 let in_dim = x.len();
1048 let out_dim = out.len();
1049 debug_assert!(w.len() >= out_dim * in_dim);
1050
1051 let row_dot = |o: usize| -> f32 {
1052 let row = &w[o * in_dim..(o + 1) * in_dim];
1053 let mut sum = 0.0f32;
1054 for j in 0..in_dim {
1055 sum += row[j] * x[j];
1056 }
1057 sum
1058 };
1059
1060 let out_addr = SendMut(out.as_mut_ptr());
1061 let run_range = move |start: usize, end: usize| {
1062 let mut o = start;
1063 // Four independent reduction chains hide add latency and reuse x.
1064 // Each row still sums j=0..in_dim in exactly the scalar order: no
1065 // horizontal SIMD reduction, FMA, or quantization approximation.
1066 while end - o >= 4 {
1067 let base = o * in_dim;
1068 let w0 = &w[base..base + in_dim];
1069 let w1 = &w[base + in_dim..base + 2 * in_dim];
1070 let w2 = &w[base + 2 * in_dim..base + 3 * in_dim];
1071 let w3 = &w[base + 3 * in_dim..base + 4 * in_dim];
1072 let (mut a, mut b, mut c, mut d) = (0.0f32, 0.0f32, 0.0f32, 0.0f32);
1073 for j in 0..in_dim {
1074 let v = x[j];
1075 a += w0[j] * v;
1076 b += w1[j] * v;
1077 c += w2[j] * v;
1078 d += w3[j] * v;
1079 }
1080 // run_rows assigns disjoint ranges within 0..out_dim.
1081 unsafe {
1082 *out_addr.at(o) = a;
1083 *out_addr.at(o + 1) = b;
1084 *out_addr.at(o + 2) = c;
1085 *out_addr.at(o + 3) = d;
1086 }
1087 o += 4;
1088 }
1089 for o in o..end {
1090 unsafe { *out_addr.at(o) = row_dot(o) };
1091 }
1092 };
1093 match pool {
1094 Some(pool) if out_dim >= 256 => pool.run_rows(out_dim, &run_range),
1095 _ => run_range(0, out_dim),
1096 }
1097}
1098
1099/// Two-input row matvec: one pass over the weight rows serves BOTH
1100/// inputs — CPU decode is memory-bound, so the second position costs a
1101/// fraction of the first (this is where MTP speculative verify wins).
1102/// Per-output accumulation order matches the single-input path exactly
1103/// → bit-identical results.
1104pub fn matvec_rows2(
1105 pool: Option<&Pool>,
1106 w: &[f32],
1107 x1: &[f32],
1108 x2: &[f32],
1109 out1: &mut [f32],
1110 out2: &mut [f32],
1111) {
1112 let in_dim = x1.len();
1113 debug_assert_eq!(x2.len(), in_dim);
1114 let out_dim = out1.len();
1115 debug_assert_eq!(out2.len(), out_dim);
1116 debug_assert!(w.len() >= out_dim * in_dim);
1117
1118 let row_dots = |o: usize| -> (f32, f32) {
1119 let row = &w[o * in_dim..(o + 1) * in_dim];
1120 let (mut s1, mut s2) = (0.0f32, 0.0f32);
1121 for j in 0..in_dim {
1122 s1 += row[j] * x1[j];
1123 s2 += row[j] * x2[j];
1124 }
1125 (s1, s2)
1126 };
1127
1128 match pool {
1129 Some(pool) if out_dim >= 256 => {
1130 let o1 = SendMut(out1.as_mut_ptr());
1131 let o2 = SendMut(out2.as_mut_ptr());
1132 let run_range = move |start: usize, end: usize| {
1133 for o in start..end {
1134 let (s1, s2) = row_dots(o);
1135 unsafe {
1136 *o1.at(o) = s1;
1137 *o2.at(o) = s2;
1138 }
1139 }
1140 };
1141 pool.run_rows(out_dim, &run_range);
1142 }
1143 _ => {
1144 for o in 0..out_dim {
1145 let (s1, s2) = row_dots(o);
1146 out1[o] = s1;
1147 out2[o] = s2;
1148 }
1149 }
1150 }
1151}
1152
1153/// `SendMut` for any element type — the sampler's sparse chain writes
1154/// per-grain candidate lists.
1155pub(crate) struct SendMutT<T>(*mut T);
1156unsafe impl<T> Send for SendMutT<T> {}
1157unsafe impl<T> Sync for SendMutT<T> {}
1158impl<T> Clone for SendMutT<T> {
1159 fn clone(&self) -> Self {
1160 *self
1161 }
1162}
1163impl<T> Copy for SendMutT<T> {}
1164impl<T> SendMutT<T> {
1165 #[inline]
1166 pub(crate) fn new(p: *mut T) -> Self {
1167 Self(p)
1168 }
1169 /// Same contract as `SendMut::at`: disjoint indices, pointee outlives
1170 /// the joined dispatch.
1171 #[inline]
1172 pub(crate) fn at(self, i: usize) -> *mut T {
1173 unsafe { self.0.add(i) }
1174 }
1175}
1176
1177#[derive(Clone, Copy)]
1178pub(crate) struct SendMut(*mut f32);
1179unsafe impl Send for SendMut {}
1180unsafe impl Sync for SendMut {}
1181
1182impl SendMut {
1183 /// The caller promises the threads it hands this to write disjoint
1184 /// indices, and that the pointee outlives them.
1185 #[inline]
1186 pub(crate) fn new(p: *mut f32) -> Self {
1187 Self(p)
1188 }
1189
1190 /// Method receiver forces the closure to capture the whole (Sync)
1191 /// wrapper, not the bare `*mut f32` field (edition-2021 precise capture).
1192 #[inline]
1193 pub(crate) fn at(self, i: usize) -> *mut f32 {
1194 unsafe { self.0.add(i) }
1195 }
1196}
1197
1198#[cfg(test)]
1199mod tests {
1200 /// A worker NOT invited to a limited job is not waited for; if it is
1201 /// preempted around the moment the job is published, the caller may
1202 /// already be on the next job. The first protocol then ran that next
1203 /// job twice and wrapped `remaining` — the caller spun forever.
1204 /// Oversubscribe the machine with spinning workers, hammer limited
1205 /// dispatches, and count every closure entry: each job must run
1206 /// exactly (limit + 1) times, and the loop must finish (a watchdog
1207 /// turns a hang into a failure). On a 2–4-CPU box the second protocol
1208 /// ran ~100 dispatches/s here (every spinner woke for every job and
1209 /// held the CPUs the invited worker needed) and hit the watchdog at
1210 /// half the jobs — slow, not wrong, but a pool that needs a time
1211 /// slice per dispatch is broken too.
1212 #[test]
1213 fn uninvited_straggler_never_runs_a_job_twice_or_wraps_the_barrier() {
1214 use std::sync::atomic::AtomicUsize;
1215 let cores = std::thread::available_parallelism().map(|n| n.get()).unwrap_or(4);
1216 let workers = (cores * 4).clamp(8, 64);
1217 let pool = Arc::new(Pool::with_spin(workers, 1_000_000));
1218 let entries = Arc::new(AtomicUsize::new(0));
1219 let expected = Arc::new(AtomicUsize::new(0));
1220 let progress = Arc::new(AtomicUsize::new(0));
1221 let done = Arc::new(AtomicBool::new(false));
1222 let iters = 20_000usize;
1223 let p = pool.clone();
1224 let (en, ex, pr, dn) = (
1225 entries.clone(),
1226 expected.clone(),
1227 progress.clone(),
1228 done.clone(),
1229 );
1230 let driver = std::thread::spawn(move || {
1231 for i in 0..iters {
1232 pr.store(i, Ordering::Relaxed);
1233 // Few grains → `run_limited` with limit < workers: most
1234 // workers are uninvited. Alternate with a full `run`.
1235 let limit = 1 + i % 3;
1236 let f = |_w: usize, _n: usize| {
1237 en.fetch_add(1, Ordering::Relaxed);
1238 };
1239 p.run_limited(limit, &f);
1240 ex.fetch_add(limit.min(p.n_workers()) + 1, Ordering::Relaxed);
1241 if i % 97 == 0 {
1242 p.run(&f);
1243 ex.fetch_add(p.n_workers() + 1, Ordering::Relaxed);
1244 }
1245 }
1246 dn.store(true, Ordering::SeqCst);
1247 });
1248 let t0 = std::time::Instant::now();
1249 while !done.load(Ordering::SeqCst) {
1250 assert!(
1251 t0.elapsed() < std::time::Duration::from_secs(120),
1252 "pool dispatch loop did not finish: at iteration {} of {}, entries {} \
1253 expected {} (the difference is the job in flight), remaining {}",
1254 progress.load(Ordering::Relaxed),
1255 iters,
1256 entries.load(Ordering::Relaxed),
1257 expected.load(Ordering::Relaxed),
1258 pool.inner.remaining.0.load(Ordering::Relaxed)
1259 );
1260 std::thread::sleep(std::time::Duration::from_millis(20));
1261 }
1262 driver.join().unwrap();
1263 // A late straggler could still be inside its (single) job; one
1264 // more full barrier drains it.
1265 pool.run(&|_w, _n| {});
1266 assert_eq!(
1267 entries.load(Ordering::Relaxed),
1268 expected.load(Ordering::Relaxed),
1269 "some job ran a closure more or fewer times than its participants"
1270 );
1271 }
1272
1273 /// Participation is decided by the ticket alone: a limited job hands
1274 /// one ticket to each of workers `0..limit` and to nobody else, a full
1275 /// `run` one to every worker. Spin 0 also drives the park/unpark path.
1276 #[test]
1277 fn limited_dispatch_hands_tickets_only_to_invited_workers() {
1278 let pool = Pool::with_spin(6, 0);
1279 let hits: Vec<AtomicUsize> = (0..7).map(|_| AtomicUsize::new(0)).collect();
1280 let bad = AtomicUsize::new(0);
1281 // No panics inside the job: a worker that dies never decrements.
1282 let f = |w: usize, n: usize| match hits.get(w) {
1283 Some(h) if w < n && (n == 3 || n == 7) => {
1284 h.fetch_add(1, Ordering::Relaxed);
1285 }
1286 _ => {
1287 bad.fetch_add(1, Ordering::Relaxed);
1288 }
1289 };
1290 let tickets = |p: &Pool| -> Vec<usize> {
1291 p.inner
1292 .slots
1293 .iter()
1294 .map(|s| s.ticket.load(Ordering::Acquire))
1295 .collect()
1296 };
1297 pool.run_limited(2, &f);
1298 assert_eq!(tickets(&pool), [1, 1, 0, 0, 0, 0]);
1299 pool.run(&f);
1300 assert_eq!(tickets(&pool), [2, 2, 1, 1, 1, 1]);
1301 assert_eq!(
1302 bad.load(Ordering::Relaxed),
1303 0,
1304 "participant index or count off"
1305 );
1306 let hits: Vec<usize> = hits.iter().map(|h| h.load(Ordering::Relaxed)).collect();
1307 // Index 2 is the caller of the limited job and worker 2 of the run.
1308 assert_eq!(hits, [2, 2, 2, 1, 1, 1, 1]);
1309 }
1310
1311 #[test]
1312 fn f32_four_row_matvec_matches_scalar_bits_and_preserves_tail() {
1313 let pool = super::Pool::new(3);
1314 for rows in [0, 1, 3, 4, 7, 255, 256, 259, 1024] {
1315 for cols in [0, 1, 3, 32, 65, 384] {
1316 let w: Vec<f32> = (0..rows * cols)
1317 .map(|i| (i as f32 * 0.173).sin() * [1e-3, 1.0, 1e3][i % 3])
1318 .collect();
1319 let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.41).cos()).collect();
1320 let want: Vec<u32> = (0..rows)
1321 .map(|r| {
1322 let mut sum = 0.0f32;
1323 for j in 0..cols {
1324 sum += w[r * cols + j] * x[j];
1325 }
1326 sum.to_bits()
1327 })
1328 .collect();
1329 for workers in [None, Some(&pool)] {
1330 let mut out = vec![17.0f32; rows + 5];
1331 super::matvec_rows(workers, &w, &x, &mut out[..rows]);
1332 let bits: Vec<u32> = out[..rows].iter().map(|v| v.to_bits()).collect();
1333 assert_eq!(bits, want, "shape {rows}x{cols}");
1334 assert_eq!(&out[rows..], &[17.0; 5]);
1335 }
1336 // A short output is an intentional public matvec contract.
1337 let mut prefix = vec![0.0f32; rows / 2];
1338 super::matvec_rows(None, &w, &x, &mut prefix);
1339 assert_eq!(
1340 prefix.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
1341 want[..rows / 2]
1342 );
1343 }
1344 }
1345 }
1346
1347 #[test]
1348 #[cfg(target_os = "linux")]
1349 fn numa_cpulist_parses_ranges_and_singles() {
1350 assert_eq!(super::numa::parse_list("0-3,8,10-11\n"), vec![0, 1, 2, 3, 8, 10, 11]);
1351 assert_eq!(super::numa::parse_list(""), Vec::<usize>::new());
1352 }
1353
1354 #[test]
1355 #[cfg(any(target_os = "android", target_os = "linux"))]
1356 fn worker_tids_registered_before_new_returns() {
1357 // WORKER_TIDS is only the last completed pool's process-wide
1358 // snapshot; another test can publish a different valid snapshot
1359 // immediately after `new` returns. Check this pool's private
1360 // registration state instead.
1361 use std::collections::HashSet;
1362 let p = super::Pool::new(3);
1363 let local: Vec<_> = p.inner.worker_tids.lock().unwrap().clone();
1364 let registered = p.inner.registered.load(Ordering::Acquire);
1365 let unique: HashSet<_> = local.iter().copied().collect();
1366 assert!(
1367 registered == 3
1368 && local.len() == 3
1369 && unique.len() == 3
1370 && local.iter().all(|&tid| tid > 0),
1371 "all worker tids must be privately registered before new returns \
1372 (registered {registered}, local {}, unique {})",
1373 local.len(),
1374 unique.len()
1375 );
1376 }
1377
1378 #[test]
1379 fn forced_threads_overrides_env_and_topology() {
1380 use std::sync::atomic::Ordering;
1381 super::FORCED_THREADS.store(3, Ordering::Relaxed);
1382 let pool = super::Pool::from_env().expect("forced 3 → pool");
1383 assert_eq!(pool.n_workers(), 3);
1384 super::FORCED_THREADS.store(1, Ordering::Relaxed);
1385 assert!(super::Pool::from_env().is_none(), "forced 1 → serial");
1386 super::FORCED_THREADS.store(0, Ordering::Relaxed);
1387 }
1388
1389 #[test]
1390 #[cfg(any(target_os = "android", target_os = "linux"))]
1391 fn concurrent_pool_constructors_complete_without_registration_race() {
1392 // WORKER_TIDS is a process-wide publication target. Before the
1393 // per-pool counter, a larger constructor could have all its workers
1394 // append, then a concurrent one could clear that vector; the larger
1395 // constructor would wait forever for a length that could never return.
1396 // Start unlike-sized constructors together so that regression is
1397 // exercised without relying on the test harness' scheduling.
1398 use std::sync::{Barrier, mpsc};
1399 use std::time::Duration;
1400
1401 for round in 0..16 {
1402 let start = Arc::new(Barrier::new(3));
1403 let (done_tx, done_rx) = mpsc::channel();
1404 let mut joins = Vec::new();
1405 for workers in [8usize, 1usize] {
1406 let start = start.clone();
1407 let done_tx = done_tx.clone();
1408 joins.push(std::thread::spawn(move || {
1409 start.wait();
1410 let pool = Pool::with_spin(workers, 0);
1411 done_tx.send(pool.n_workers()).unwrap();
1412 }));
1413 }
1414 drop(done_tx);
1415 start.wait();
1416 let mut sizes = Vec::with_capacity(2);
1417 for _ in 0..2 {
1418 sizes.push(
1419 done_rx
1420 .recv_timeout(Duration::from_secs(10))
1421 .unwrap_or_else(|_| panic!("pool constructor stalled in round {round}")),
1422 );
1423 }
1424 sizes.sort_unstable();
1425 assert_eq!(sizes, [1, 8]);
1426 for join in joins {
1427 join.join().unwrap();
1428 }
1429 }
1430 }
1431
1432 #[test]
1433 fn capacity_split_clock_bins_vs_microarch() {
1434 type P = super::Pool;
1435 // JR510: all-A55, two clock bins — use every core.
1436 assert_eq!(
1437 P::cores_from_capacities(&[1024, 1024, 1024, 1024, 768, 768, 768, 768]),
1438 Some(8)
1439 );
1440 // Classic big.LITTLE (A78 + A55) — big only.
1441 assert_eq!(
1442 P::cores_from_capacities(&[1024, 1024, 1024, 1024, 350, 350, 350, 350]),
1443 Some(4)
1444 );
1445 // Three-tier flagship: X + A7xx mids stay, A5xx littles go.
1446 assert_eq!(
1447 P::cores_from_capacities(&[1024, 800, 800, 800, 800, 300, 300, 300]),
1448 Some(5)
1449 );
1450 // Uniform: no signal, caller falls back.
1451 assert_eq!(P::cores_from_capacities(&[1024; 8]), None);
1452 assert_eq!(P::cores_from_capacities(&[]), None);
1453 }
1454
1455 use super::*;
1456
1457 #[test]
1458 fn parallel_matvec_equals_serial_bitexact() {
1459 let (out_dim, in_dim) = (512, 64);
1460 let w: Vec<f32> = (0..out_dim * in_dim)
1461 .map(|i| (i as f32 * 0.013).sin())
1462 .collect();
1463 let x: Vec<f32> = (0..in_dim).map(|i| (i as f32 * 0.07).cos()).collect();
1464
1465 let mut serial = vec![0.0f32; out_dim];
1466 matvec_rows(None, &w, &x, &mut serial);
1467
1468 let pool = Pool::new(4);
1469 let mut parallel = vec![0.0f32; out_dim];
1470 matvec_rows(Some(&pool), &w, &x, &mut parallel);
1471
1472 assert_eq!(serial, parallel, "row-parallel must be bit-identical");
1473 }
1474
1475 #[test]
1476 fn fused_pair_equals_two_singles_bitexact() {
1477 let (out_dim, in_dim) = (300, 48);
1478 let w: Vec<f32> = (0..out_dim * in_dim)
1479 .map(|i| (i as f32 * 0.011).sin())
1480 .collect();
1481 let x1: Vec<f32> = (0..in_dim).map(|i| (i as f32 * 0.03).cos()).collect();
1482 let x2: Vec<f32> = (0..in_dim).map(|i| (i as f32 * 0.09).sin()).collect();
1483
1484 let mut a1 = vec![0.0f32; out_dim];
1485 let mut a2 = vec![0.0f32; out_dim];
1486 matvec_rows(None, &w, &x1, &mut a1);
1487 matvec_rows(None, &w, &x2, &mut a2);
1488
1489 for pool in [None, Some(Pool::new(3))] {
1490 let mut b1 = vec![0.0f32; out_dim];
1491 let mut b2 = vec![0.0f32; out_dim];
1492 matvec_rows2(pool.as_ref(), &w, &x1, &x2, &mut b1, &mut b2);
1493 assert_eq!(a1, b1, "fused lane 1 must be bit-identical");
1494 assert_eq!(a2, b2, "fused lane 2 must be bit-identical");
1495 }
1496 }
1497
1498 #[test]
1499 fn pool_survives_many_runs() {
1500 let pool = Pool::new(3);
1501 let counter = AtomicUsize::new(0);
1502 for _ in 0..100 {
1503 pool.run(&|_, _| {
1504 counter.fetch_add(1, Ordering::Relaxed);
1505 });
1506 }
1507 // 3 workers + the participating caller = 4 executions per run.
1508 assert_eq!(counter.load(Ordering::Relaxed), 400);
1509 }
1510
1511 #[test]
1512 fn pool_wakes_after_park() {
1513 // Force immediate parking (no spin) — the epoch/parked handshake
1514 // must still never miss a wakeup.
1515 let pool = Pool::with_spin(2, 0);
1516 let counter = AtomicUsize::new(0);
1517 for _ in 0..50 {
1518 pool.run(&|_, _| {
1519 counter.fetch_add(1, Ordering::Relaxed);
1520 });
1521 // Give workers time to actually park between jobs.
1522 std::thread::sleep(std::time::Duration::from_micros(200));
1523 }
1524 assert_eq!(counter.load(Ordering::Relaxed), 150);
1525 }
1526
1527 #[test]
1528 fn worker_indices_are_distinct_and_cover_range() {
1529 let pool = Pool::new(3);
1530 let hits: Vec<AtomicUsize> = (0..4).map(|_| AtomicUsize::new(0)).collect();
1531 for _ in 0..20 {
1532 pool.run(&|widx, n| {
1533 assert_eq!(n, 4);
1534 hits[widx].fetch_add(1, Ordering::Relaxed);
1535 });
1536 }
1537 for (i, h) in hits.iter().enumerate() {
1538 assert_eq!(h.load(Ordering::Relaxed), 20, "participant {i} missed runs");
1539 }
1540 }
1541}
1542
1543#[cfg(test)]
1544mod grain_tests {
1545 use super::grain_for;
1546
1547 #[test]
1548 fn a_short_job_still_reaches_every_worker() {
1549 // 24 rows, 49 workers: the old flat floor of 32 handed all 24 to the
1550 // first worker and woke the rest for nothing.
1551 assert_eq!(grain_for(24, 49), 1);
1552 // Wide jobs keep the stride the SDOT loop wants.
1553 assert_eq!(grain_for(4096, 49), 32);
1554 assert_eq!(grain_for(32768, 49), 83);
1555 // Degenerate shapes must not divide by zero or return zero.
1556 assert_eq!(grain_for(0, 49), 1);
1557 assert_eq!(grain_for(7, 1), 7);
1558 assert!(grain_for(1, 49) >= 1);
1559 }
1560}