use std::cell::RefCell;
#[cfg(feature = "multithread-mm")]
use std::sync::atomic::{AtomicUsize, Ordering};
#[allow(unused_imports)]
use std::sync::{Arc, Mutex};
#[cfg(feature = "multithread-mm")]
use rayon::{ThreadPool, ThreadPoolBuilder};
#[cfg(feature = "multithread-mm")]
use tract_data::internal::vector_size;
use tract_data::internal::{Tensor, TensorView, TractResult, ensure};
use crate::LinalgFn;
#[derive(Debug, Clone, Default)]
pub enum Executor {
#[default]
SingleThread,
#[cfg(feature = "multithread-mm")]
MultiThread(Arc<ThreadPool>),
#[cfg(feature = "multithread-mm")]
RayonGlobal,
}
impl Executor {
#[cfg(feature = "multithread-mm")]
pub fn multithread(n: usize) -> Executor {
Executor::multithread_with_name(n, "tract-default")
}
#[cfg(feature = "multithread-mm")]
pub fn multithread_with_name(n: usize, name: &str) -> Executor {
let name = name.to_string();
let pool = ThreadPoolBuilder::new()
.thread_name(move |n| format!("{name}-{n}"))
.num_threads(n)
.build()
.unwrap();
Executor::MultiThread(Arc::new(pool))
}
}
static DEFAULT_EXECUTOR: Mutex<Executor> = Mutex::new(Executor::SingleThread);
thread_local! {
static TLS_EXECUTOR_OVERRIDE: RefCell<Option<Executor>> = Default::default();
}
pub fn current_tract_executor() -> Executor {
if let Some(over_ride) = TLS_EXECUTOR_OVERRIDE.with_borrow(|tls| tls.clone()) {
over_ride
} else {
DEFAULT_EXECUTOR.lock().unwrap().clone()
}
}
pub fn set_default_executor(executor: Executor) {
*DEFAULT_EXECUTOR.lock().unwrap() = executor;
}
pub fn multithread_tract_scope<R, F: FnOnce() -> R>(pool: Executor, f: F) -> R {
let previous = TLS_EXECUTOR_OVERRIDE.replace(Some(pool));
let result = f();
TLS_EXECUTOR_OVERRIDE.set(previous);
result
}
#[cfg(feature = "multithread-mm")]
static THREADING_PANEL_THRESHOLD: AtomicUsize = AtomicUsize::new(64);
#[cfg(feature = "multithread-mm")]
pub fn current_threading_panel_threshold() -> usize {
THREADING_PANEL_THRESHOLD.load(Ordering::Relaxed)
}
#[cfg(feature = "multithread-mm")]
pub fn set_threading_panel_threshold(panels: usize) {
THREADING_PANEL_THRESHOLD.store(panels, Ordering::Relaxed);
}
#[cfg(feature = "multithread-mm")]
static THREADING_ELEMENT_THRESHOLD: AtomicUsize = AtomicUsize::new(32768);
#[cfg(feature = "multithread-mm")]
pub fn current_threading_element_threshold() -> usize {
THREADING_ELEMENT_THRESHOLD.load(Ordering::Relaxed)
}
#[cfg(feature = "multithread-mm")]
pub fn set_threading_element_threshold(elements: usize) {
THREADING_ELEMENT_THRESHOLD.store(elements, Ordering::Relaxed);
}
pub fn par_chunks_mut<T: Send>(
out: &mut [T],
row_len: usize,
total_elems: usize,
f: impl Fn(usize, &mut [T]) -> TractResult<()> + Sync + Send,
) -> TractResult<()> {
#[cfg(feature = "multithread-mm")]
{
use rayon::prelude::*;
debug_assert!(row_len >= 1 && out.len() % row_len == 0);
let n_rows = out.len() / row_len;
if n_rows < 2 || total_elems < current_threading_element_threshold() {
return f(0, out);
}
let run = |out: &mut [T]| -> TractResult<()> {
let n_chunks = (4 * rayon::current_num_threads()).min(n_rows);
let chunk_rows = n_rows.div_ceil(n_chunks);
out.par_chunks_mut(chunk_rows * row_len)
.enumerate()
.try_for_each(|(i, chunk)| f(i * chunk_rows, chunk))
};
match current_tract_executor() {
Executor::MultiThread(pool) if pool.current_num_threads() > 1 => {
pool.install(|| run(out))
}
Executor::RayonGlobal => run(out),
_ => f(0, out),
}
}
#[cfg(not(feature = "multithread-mm"))]
{
let _ = (row_len, total_elems);
f(0, out)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum BShare {
Lockstep,
PerBlock,
}
pub fn par_bin(
eval_fn: &LinalgFn,
a: &mut Tensor,
b: &Tensor,
period: usize,
share: BShare,
) -> TractResult<()> {
if a.len() == 0 {
return Ok(());
}
ensure!(
a.len().is_multiple_of(period),
"par_bin: period {period} does not divide a.len() {}",
a.len()
);
let n_blocks = a.len() / period;
match share {
BShare::Lockstep => ensure!(
b.len() == period,
"par_bin: Lockstep wants b.len() {} == period {period}",
b.len()
),
BShare::PerBlock => ensure!(
b.len() == n_blocks,
"par_bin: PerBlock wants b.len() {} == {n_blocks} blocks",
b.len()
),
}
let a_item = a.datum_type().size_of() as isize;
let b_item = b.datum_type().size_of() as isize;
let a = &*a;
let call = |block: usize, offset: usize, len: usize| -> TractResult<()> {
static STRIDES: [isize; 1] = [1];
let a_shape = [len];
let (b_offset, b_shape) = match share {
BShare::Lockstep => (offset as isize * b_item, [len]),
BShare::PerBlock => (block as isize * b_item, [1]),
};
let a_offset = (block * period + offset) as isize * a_item;
unsafe {
let mut a_chunk = TensorView::from_bytes(a, a_offset, &a_shape, &STRIDES);
let b_chunk = TensorView::from_bytes(b, b_offset, &b_shape, &STRIDES);
eval_fn(&mut a_chunk, &b_chunk)
}
};
#[cfg(feature = "multithread-mm")]
{
use rayon::prelude::*;
if a.len() >= current_threading_element_threshold() {
let executor = current_tract_executor();
let nth = match &executor {
Executor::MultiThread(pool) => pool.current_num_threads(),
Executor::RayonGlobal => rayon::current_num_threads(),
Executor::SingleThread => 1,
};
if nth > 1 {
let per_block = (4 * nth).div_ceil(n_blocks).max(1);
let chunk = period.div_ceil(per_block).next_multiple_of(vector_size()).min(period);
let per_block = period.div_ceil(chunk);
if n_blocks * per_block > 1 {
let run = || {
(0..n_blocks * per_block).into_par_iter().try_for_each(|i| {
let offset = (i % per_block) * chunk;
call(i / per_block, offset, chunk.min(period - offset))
})
};
return match executor {
Executor::MultiThread(pool) => pool.install(run),
_ => run(),
};
}
}
}
}
if n_blocks == 1 {
return eval_fn(&mut a.view(), &b.view());
}
(0..n_blocks).try_for_each(|block| call(block, 0, period))
}