Skip to main content

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}