use cubecl_common::device::ServiceId;
use cubecl_environment::{collections::HashMap, sync::Mutex};
use cubecl_runtime::{client::Client, throughput::MemoryAccess};
use crate::throughput::LaunchConfig;
const MEMORY_WORKER_SHAPES: usize = 5;
static SATURATING_WORKERS: Mutex<Option<HashMap<(ServiceId, MemoryAccess), u32>>> =
Mutex::new(None);
pub(super) struct WorkerSweep;
impl WorkerSweep {
pub(super) fn shapes(
client: &Client,
config: LaunchConfig,
access: MemoryAccess,
) -> alloc::vec::Vec<LaunchConfig> {
if client.properties().hardware.num_cpu_cores.is_none() {
return alloc::vec![config];
}
if let Some(units) = Self::remembered(client, access) {
return alloc::vec![config.with_units(client, units)];
}
Self::counts(config.cube_dim.num_elems() as usize)
.into_iter()
.map(|units| config.with_units(client, units as u32))
.collect()
}
fn counts(units: usize) -> alloc::vec::Vec<usize> {
core::iter::successors(Some(units.max(1)), |units| (*units > 1).then(|| units / 2))
.take(MEMORY_WORKER_SHAPES)
.collect()
}
fn remembered(client: &Client, access: MemoryAccess) -> Option<u32> {
let workers = SATURATING_WORKERS.lock();
let workers = workers.as_ref()?;
workers.get(&(client.service_id(), access)).copied()
}
pub(super) fn remember(client: &Client, access: MemoryAccess, units: u32) {
let mut workers = SATURATING_WORKERS.lock();
workers
.get_or_insert_with(HashMap::new)
.insert((client.service_id(), access), units);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn the_sweep_starts_at_the_full_launch_and_halves_to_the_budget() {
assert_eq!(WorkerSweep::counts(16), alloc::vec![16, 8, 4, 2, 1]);
assert_eq!(WorkerSweep::counts(128), alloc::vec![128, 64, 32, 16, 8]);
}
#[test]
fn one_core_is_still_a_shape() {
assert_eq!(WorkerSweep::counts(1), alloc::vec![1]);
}
}