use crate::error::WhisperResult;
#[cfg(feature = "parallel")]
pub fn configure_thread_pool(num_threads: Option<u32>) -> WhisperResult<usize> {
use rayon::ThreadPoolBuilder;
let builder = ThreadPoolBuilder::new();
let builder = if let Some(n) = num_threads {
builder.num_threads(n as usize)
} else {
let logical = std::thread::available_parallelism()
.map(|n| n.get() as u32)
.unwrap_or(4);
let physical = (logical / 2).clamp(1, 16);
builder.num_threads(physical as usize)
};
match builder.build_global() {
Ok(()) => Ok(rayon::current_num_threads()),
Err(_) => {
Ok(rayon::current_num_threads())
}
}
}
#[cfg(not(feature = "parallel"))]
pub fn configure_thread_pool(num_threads: Option<u32>) -> WhisperResult<usize> {
let _ = num_threads;
Ok(1)
}
#[cfg(feature = "parallel")]
pub fn parallel_map<T, F>(range: std::ops::Range<usize>, f: F) -> Vec<T>
where
T: Send,
F: Fn(usize) -> T + Send + Sync,
{
use rayon::prelude::*;
range.into_par_iter().map(f).collect()
}
#[cfg(not(feature = "parallel"))]
pub fn parallel_map<T, F>(range: std::ops::Range<usize>, f: F) -> Vec<T>
where
F: Fn(usize) -> T,
{
range.map(f).collect()
}
#[cfg(feature = "parallel")]
pub fn parallel_try_map<T, F>(range: std::ops::Range<usize>, f: F) -> WhisperResult<Vec<T>>
where
T: Send,
F: Fn(usize) -> WhisperResult<T> + Send + Sync,
{
use rayon::prelude::*;
range.into_par_iter().map(f).collect()
}
#[cfg(not(feature = "parallel"))]
pub fn parallel_try_map<T, F>(range: std::ops::Range<usize>, f: F) -> WhisperResult<Vec<T>>
where
F: Fn(usize) -> WhisperResult<T>,
{
range.map(f).collect()
}
#[cfg(feature = "parallel")]
pub fn is_parallel_available() -> bool {
#[cfg(target_arch = "wasm32")]
{
crate::wasm::threading::is_threaded_available()
}
#[cfg(not(target_arch = "wasm32"))]
{
true
}
}
#[cfg(not(feature = "parallel"))]
pub fn is_parallel_available() -> bool {
false
}
#[cfg(feature = "parallel")]
pub fn thread_count() -> usize {
#[cfg(target_arch = "wasm32")]
{
crate::wasm::threading::optimal_thread_count()
}
#[cfg(not(target_arch = "wasm32"))]
{
rayon::current_num_threads()
}
}
#[cfg(not(feature = "parallel"))]
pub fn thread_count() -> usize {
1
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_parallel_map_basic() {
let results = parallel_map(0..4, |i| i * 2);
assert_eq!(results, vec![0, 2, 4, 6]);
}
#[test]
fn test_parallel_map_order_preserved() {
let results = parallel_map(0..100, |i| i);
let expected: Vec<usize> = (0..100).collect();
assert_eq!(results, expected);
}
#[test]
fn test_parallel_try_map_success() {
let results = parallel_try_map(0..4, |i| Ok(i * 2));
assert!(results.is_ok());
assert_eq!(
results.expect("parallel_try_map should succeed"),
vec![0, 2, 4, 6]
);
}
#[test]
fn test_parallel_try_map_error() {
let results: WhisperResult<Vec<i32>> = parallel_try_map(0..4, |i| {
if i == 2 {
Err(crate::error::WhisperError::Model("test error".into()))
} else {
Ok(i as i32)
}
});
assert!(results.is_err());
}
#[test]
fn test_thread_count() {
let count = thread_count();
assert!(count >= 1);
#[cfg(feature = "parallel")]
{
println!("Thread count: {}", count);
}
}
#[test]
fn test_is_parallel_available() {
let available = is_parallel_available();
#[cfg(feature = "parallel")]
assert!(available, "parallel feature enabled but not available");
#[cfg(not(feature = "parallel"))]
assert!(
!available,
"parallel should not be available without feature"
);
}
#[test]
fn test_configure_thread_pool() {
let result = configure_thread_pool(Some(2));
assert!(result.is_ok());
let threads = result.expect("configure_thread_pool should succeed");
assert!(threads >= 1);
}
#[test]
fn test_configure_thread_pool_default() {
let result = configure_thread_pool(None);
assert!(result.is_ok());
}
#[test]
fn test_parallel_map_empty() {
let results: Vec<i32> = parallel_map(0..0, |i| i as i32);
assert!(results.is_empty());
}
#[test]
fn test_parallel_map_single() {
let results = parallel_map(0..1, |i| i * 10);
assert_eq!(results, vec![0]);
}
#[test]
fn test_parallel_try_map_empty() {
let results: WhisperResult<Vec<i32>> = parallel_try_map(0..0, |i| Ok(i as i32));
assert!(results.is_ok());
assert!(results
.expect("parallel_try_map should succeed on empty")
.is_empty());
}
}