use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::mpsc::{channel, Sender};
use std::sync::{Arc, Condvar, Mutex};
struct TaskPtr(*const (dyn Fn(usize, usize) + Sync));
unsafe impl Send for TaskPtr {}
struct Latch {
remaining: AtomicUsize,
lock: Mutex<()>,
cv: Condvar,
}
impl Latch {
fn new(n: usize) -> Arc<Self> {
Arc::new(Self {
remaining: AtomicUsize::new(n),
lock: Mutex::new(()),
cv: Condvar::new(),
})
}
fn count_down(&self) {
if self.remaining.fetch_sub(1, Ordering::AcqRel) == 1 {
let _g = self.lock.lock().unwrap();
self.cv.notify_all();
}
}
fn wait(&self) {
let mut g = self.lock.lock().unwrap();
while self.remaining.load(Ordering::Acquire) != 0 {
g = self.cv.wait(g).unwrap();
}
}
}
struct Job {
task: TaskPtr,
worker_idx: usize,
n_workers: usize,
latch: Arc<Latch>,
}
pub struct Pool {
txs: Vec<Sender<Job>>,
}
impl Pool {
pub fn new(n_workers: usize) -> Self {
let mut txs = Vec::with_capacity(n_workers);
for w in 0..n_workers {
let (tx, rx) = channel::<Job>();
std::thread::Builder::new()
.name(format!("cmf-pool-{w}"))
.spawn(move || {
while let Ok(job) = rx.recv() {
let f = unsafe { &*job.task.0 };
f(job.worker_idx, job.n_workers);
job.latch.count_down();
}
})
.expect("spawn pool worker");
txs.push(tx);
}
Self { txs }
}
pub fn from_env() -> Option<Arc<Self>> {
let n = match std::env::var("CMF_THREADS") {
Ok(v) => v.parse::<usize>().unwrap_or(0),
Err(_) => {
let avail = std::thread::available_parallelism()
.map(|n| n.get())
.unwrap_or(1);
avail.saturating_sub(1).min(8)
}
};
if n <= 1 {
None
} else {
Some(Arc::new(Self::new(n)))
}
}
pub fn n_workers(&self) -> usize {
self.txs.len()
}
pub fn run(&self, f: &(dyn Fn(usize, usize) + Sync)) {
let n = self.txs.len();
let latch = Latch::new(n);
let ptr: *const (dyn Fn(usize, usize) + Sync) = f;
let ptr: *const (dyn Fn(usize, usize) + Sync + 'static) =
unsafe { std::mem::transmute(ptr) };
for (i, tx) in self.txs.iter().enumerate() {
let job = Job {
task: TaskPtr(ptr),
worker_idx: i,
n_workers: n,
latch: latch.clone(),
};
tx.send(job).expect("pool worker died");
}
latch.wait();
}
}
pub fn matvec_rows(pool: Option<&Pool>, w: &[f32], x: &[f32], out: &mut [f32]) {
let in_dim = x.len();
let out_dim = out.len();
debug_assert!(w.len() >= out_dim * in_dim);
let row_dot = |o: usize| -> f32 {
let row = &w[o * in_dim..(o + 1) * in_dim];
let mut sum = 0.0f32;
for j in 0..in_dim {
sum += row[j] * x[j];
}
sum
};
match pool {
Some(pool) if out_dim >= 256 => {
let out_addr = SendMut(out.as_mut_ptr());
pool.run(&move |widx, n| {
let chunk = out_dim.div_ceil(n);
let start = widx * chunk;
let end = (start + chunk).min(out_dim);
for o in start..end {
unsafe { *out_addr.at(o) = row_dot(o) };
}
});
}
_ => {
for (o, dst) in out.iter_mut().enumerate() {
*dst = row_dot(o);
}
}
}
}
pub fn matvec_rows2(
pool: Option<&Pool>,
w: &[f32],
x1: &[f32],
x2: &[f32],
out1: &mut [f32],
out2: &mut [f32],
) {
let in_dim = x1.len();
debug_assert_eq!(x2.len(), in_dim);
let out_dim = out1.len();
debug_assert_eq!(out2.len(), out_dim);
debug_assert!(w.len() >= out_dim * in_dim);
let row_dots = |o: usize| -> (f32, f32) {
let row = &w[o * in_dim..(o + 1) * in_dim];
let (mut s1, mut s2) = (0.0f32, 0.0f32);
for j in 0..in_dim {
s1 += row[j] * x1[j];
s2 += row[j] * x2[j];
}
(s1, s2)
};
match pool {
Some(pool) if out_dim >= 256 => {
let o1 = SendMut(out1.as_mut_ptr());
let o2 = SendMut(out2.as_mut_ptr());
pool.run(&move |widx, n| {
let chunk = out_dim.div_ceil(n);
let start = widx * chunk;
let end = (start + chunk).min(out_dim);
for o in start..end {
let (s1, s2) = row_dots(o);
unsafe {
*o1.at(o) = s1;
*o2.at(o) = s2;
}
}
});
}
_ => {
for o in 0..out_dim {
let (s1, s2) = row_dots(o);
out1[o] = s1;
out2[o] = s2;
}
}
}
}
#[derive(Clone, Copy)]
struct SendMut(*mut f32);
unsafe impl Send for SendMut {}
unsafe impl Sync for SendMut {}
impl SendMut {
#[inline]
fn at(self, i: usize) -> *mut f32 {
unsafe { self.0.add(i) }
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parallel_matvec_equals_serial_bitexact() {
let (out_dim, in_dim) = (512, 64);
let w: Vec<f32> = (0..out_dim * in_dim).map(|i| (i as f32 * 0.013).sin()).collect();
let x: Vec<f32> = (0..in_dim).map(|i| (i as f32 * 0.07).cos()).collect();
let mut serial = vec![0.0f32; out_dim];
matvec_rows(None, &w, &x, &mut serial);
let pool = Pool::new(4);
let mut parallel = vec![0.0f32; out_dim];
matvec_rows(Some(&pool), &w, &x, &mut parallel);
assert_eq!(serial, parallel, "row-parallel must be bit-identical");
}
#[test]
fn fused_pair_equals_two_singles_bitexact() {
let (out_dim, in_dim) = (300, 48);
let w: Vec<f32> = (0..out_dim * in_dim).map(|i| (i as f32 * 0.011).sin()).collect();
let x1: Vec<f32> = (0..in_dim).map(|i| (i as f32 * 0.03).cos()).collect();
let x2: Vec<f32> = (0..in_dim).map(|i| (i as f32 * 0.09).sin()).collect();
let mut a1 = vec![0.0f32; out_dim];
let mut a2 = vec![0.0f32; out_dim];
matvec_rows(None, &w, &x1, &mut a1);
matvec_rows(None, &w, &x2, &mut a2);
for pool in [None, Some(Pool::new(3))] {
let mut b1 = vec![0.0f32; out_dim];
let mut b2 = vec![0.0f32; out_dim];
matvec_rows2(pool.as_ref(), &w, &x1, &x2, &mut b1, &mut b2);
assert_eq!(a1, b1, "fused lane 1 must be bit-identical");
assert_eq!(a2, b2, "fused lane 2 must be bit-identical");
}
}
#[test]
fn pool_survives_many_runs() {
let pool = Pool::new(3);
let counter = AtomicUsize::new(0);
for _ in 0..100 {
pool.run(&|_, _| {
counter.fetch_add(1, Ordering::Relaxed);
});
}
assert_eq!(counter.load(Ordering::Relaxed), 300);
}
}