use std::{
collections::HashMap,
sync::{Arc, Mutex, Weak},
};
use laddu_autodiff::AutodiffMode;
use laddu_data::io::{Partitioning, ReadPlan};
#[cfg(feature = "wgpu")]
use laddu_memory::{DeviceIdentity, MemoryResource};
use laddu_memory::{
MemoryBudget, MemoryDecision, MemoryPlan, MemoryPool, MemoryPoolReport, MemoryReport,
MemoryState,
};
use rayon::{ThreadPool, ThreadPoolBuilder};
use serde::{Deserialize, Serialize};
#[cfg(feature = "wgpu")]
use crate::RuntimeError;
use crate::{ExecutionError, RuntimeResult};
pub(crate) type NormalizationCache =
HashMap<(u64, u64, NormalizationMode), Weak<crate::PreparedNormalization>>;
#[cfg(feature = "mpi")]
use mpi::{
collective::SystemOperation,
topology::SimpleCommunicator,
traits::{Communicator, CommunicatorCollectives},
};
#[derive(Copy, Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
pub enum Precision {
#[default]
Auto,
F32,
F64,
}
#[derive(Copy, Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
pub enum ThreadPolicy {
#[default]
Auto,
Serial,
Fixed(usize),
}
#[derive(Copy, Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
pub enum JitPolicy {
#[default]
Auto,
Enabled,
Disabled,
}
#[derive(Copy, Clone, Debug, Default, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub enum NormalizationMode {
#[default]
Auto,
General,
Verify,
}
#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct CpuOptions {
pub threads: ThreadPolicy,
pub jit: JitPolicy,
}
#[derive(Copy, Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
pub enum GpuBackend {
#[default]
Auto,
Wgpu,
Cuda,
}
#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
pub enum GpuDeviceSelector {
#[default]
Auto,
Index(usize),
PciBusId(String),
Name(String),
}
#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct GpuOptions {
pub backend: GpuBackend,
pub device: GpuDeviceSelector,
}
#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
pub enum Device {
#[default]
Auto,
Cpu(CpuOptions),
Gpu(GpuOptions),
}
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
pub struct ExecutionOptions {
pub device: Device,
pub precision: Precision,
pub autodiff: AutodiffMode,
#[serde(default)]
pub normalization: NormalizationMode,
pub partitioning: Partitioning,
pub memory: MemoryPlan,
}
#[derive(Clone)]
pub struct Execution {
requested_device: Device,
precision: Precision,
autodiff: AutodiffMode,
normalization: NormalizationMode,
threads: ThreadPolicy,
jit: JitPolicy,
pool: Option<Arc<ThreadPool>>,
partitioning: Partitioning,
memory_state: MemoryState,
host_memory: MemoryPool,
device_memory: Option<MemoryPool>,
memory_decisions: Arc<Mutex<Vec<MemoryDecision>>>,
normalization_cache: Arc<Mutex<NormalizationCache>>,
#[cfg(feature = "wgpu")]
wgpu: Option<Arc<laddu_wgpu::WgpuContext>>,
#[cfg(feature = "mpi")]
communicator: Option<Arc<SimpleCommunicator>>,
}
impl std::fmt::Debug for Execution {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
#[cfg(feature = "wgpu")]
let resolved_device = if self.wgpu.is_some() { "wgpu" } else { "cpu" };
#[cfg(not(feature = "wgpu"))]
let resolved_device = "cpu";
formatter
.debug_struct("Execution")
.field("requested_device", &self.requested_device)
.field("resolved_device", &resolved_device)
.field("precision", &self.precision)
.field("autodiff", &self.autodiff)
.field("normalization", &self.normalization)
.field("threads", &self.threads)
.field("jit", &self.jit)
.field("partitioning", &self.partitioning)
.field("host_memory", &self.host_memory.report())
.field(
"device_memory",
&self.device_memory.as_ref().map(MemoryPool::report),
)
.field("ranks", &self.nranks())
.finish_non_exhaustive()
}
}
impl Default for Execution {
fn default() -> Self {
let memory_state = MemoryState::current();
memory_state.refresh();
let host_memory = memory_state
.pool("host", MemoryBudget::Auto)
.expect("host memory discovery must resolve an automatic budget");
Self {
requested_device: Device::Auto,
precision: Precision::F64,
autodiff: AutodiffMode::Auto,
normalization: NormalizationMode::Auto,
threads: ThreadPolicy::Auto,
jit: JitPolicy::Auto,
pool: None,
partitioning: Partitioning::default(),
memory_state,
host_memory,
device_memory: None,
memory_decisions: Default::default(),
normalization_cache: Default::default(),
#[cfg(feature = "wgpu")]
wgpu: None,
#[cfg(feature = "mpi")]
communicator: None,
}
}
}
impl Execution {
pub fn local(options: ExecutionOptions) -> RuntimeResult<Self> {
let memory_state = MemoryState::current();
memory_state.refresh();
let host_memory = memory_state.pool("host", options.memory.host)?;
#[cfg(feature = "wgpu")]
let mut wgpu = None;
#[cfg(feature = "wgpu")]
let mut device_memory = None;
#[cfg(not(feature = "wgpu"))]
let device_memory = None;
let cpu = match &options.device {
Device::Auto => CpuOptions::default(),
Device::Cpu(options) => options.clone(),
Device::Gpu(gpu_options) => {
#[cfg(feature = "wgpu")]
{
if gpu_options.backend == GpuBackend::Cuda {
return Err(ExecutionError::GpuUnavailable(gpu_options.backend).into());
}
let selector = match &gpu_options.device {
GpuDeviceSelector::Auto => laddu_wgpu::WgpuDeviceSelector::Auto,
GpuDeviceSelector::Index(index) => {
laddu_wgpu::WgpuDeviceSelector::Index(*index)
}
GpuDeviceSelector::PciBusId(id) => {
laddu_wgpu::WgpuDeviceSelector::PciBusId(id.clone())
}
GpuDeviceSelector::Name(name) => {
laddu_wgpu::WgpuDeviceSelector::Name(name.clone())
}
};
let precision = match options.precision {
Precision::Auto => laddu_wgpu::WgpuPrecision::Auto,
Precision::F32 => laddu_wgpu::WgpuPrecision::F32,
Precision::F64 => laddu_wgpu::WgpuPrecision::F64,
};
let mut context = laddu_wgpu::WgpuBackend::default()
.open(
&laddu_wgpu::WgpuOptions {
device: selector,
memory_budget: None,
},
precision,
)
.map_err(|error| RuntimeError::Wgpu(error.to_string()))?;
let resource_id = if context.info().pci_bus_id.is_empty() {
format!("wgpu:{}", context.info().index)
} else {
format!("pci:{}", context.info().pci_bus_id)
};
let fallback = context
.info()
.max_buffer_size
.min(512 * 1024 * 1024)
.max(context.info().max_storage_buffer_binding_size);
let resource = MemoryResource::discover_device(
resource_id.clone(),
context.info().name.clone(),
DeviceIdentity {
adapter_index: context.info().index,
vendor_id: context.info().vendor,
device_id: context.info().device,
pci_bus_id: context.info().pci_bus_id.clone(),
},
fallback,
);
memory_state.register_device(resource);
let requested = options.memory.device.unwrap_or(MemoryBudget::Auto);
let pool = memory_state.pool(&resource_id, requested)?;
context
.set_memory_budget(usize::try_from(pool.capacity()).unwrap_or(usize::MAX));
device_memory = Some(pool);
wgpu = Some(Arc::new(context));
CpuOptions::default()
}
#[cfg(not(feature = "wgpu"))]
return Err(ExecutionError::GpuUnavailable(gpu_options.backend).into());
}
};
let precision = match options.precision {
Precision::Auto if matches!(options.device, Device::Gpu(_)) => Precision::F32,
Precision::Auto => Precision::F64,
precision => precision,
};
#[cfg(not(feature = "jit"))]
if cpu.jit == JitPolicy::Enabled {
return Err(ExecutionError::JitUnavailable.into());
}
let pool = match cpu.threads {
ThreadPolicy::Fixed(0) => return Err(ExecutionError::ZeroThreads.into()),
ThreadPolicy::Fixed(threads) => Some(Arc::new(
ThreadPoolBuilder::new()
.num_threads(threads)
.build()
.map_err(|error| ExecutionError::ThreadPool(error.to_string()))?,
)),
ThreadPolicy::Auto | ThreadPolicy::Serial => None,
};
Ok(Self {
requested_device: options.device,
precision,
autodiff: options.autodiff,
normalization: options.normalization,
threads: cpu.threads,
jit: cpu.jit,
pool,
partitioning: options.partitioning,
memory_state,
host_memory,
device_memory,
memory_decisions: Default::default(),
normalization_cache: Default::default(),
#[cfg(feature = "wgpu")]
wgpu,
#[cfg(feature = "mpi")]
communicator: None,
})
}
#[cfg(feature = "mpi")]
pub fn distributed<C>(options: ExecutionOptions, world: &C) -> RuntimeResult<Self>
where
C: Communicator,
{
let local_processes = mpi_local_process_count(world.size());
let mut options = options;
options.memory.host = shared_mpi_budget(options.memory.host, local_processes);
options.memory.device = options
.memory
.device
.map(|budget| shared_mpi_budget(budget, local_processes));
let mut execution = Self::local(options)?;
execution.record_memory_decision(MemoryDecision {
label: "mpi-memory-share".into(),
fixed_bytes: 0,
bytes_per_event: 0,
chunk_events: 0,
estimated_peak_bytes: 0,
actual_high_water_bytes: None,
strategy: format!("equal-share-across-{local_processes}-local-ranks"),
});
execution.communicator = Some(Arc::new(world.duplicate()));
Ok(execution)
}
pub fn requested_device(&self) -> &Device {
&self.requested_device
}
#[cfg(feature = "wgpu")]
pub(crate) fn wgpu_context(&self) -> Option<&Arc<laddu_wgpu::WgpuContext>> {
self.wgpu.as_ref()
}
pub fn precision(&self) -> Precision {
self.precision
}
pub fn autodiff_mode(&self) -> AutodiffMode {
self.autodiff
}
pub fn normalization_mode(&self) -> NormalizationMode {
self.normalization
}
pub(crate) fn normalization_cache(&self) -> &Mutex<NormalizationCache> {
&self.normalization_cache
}
pub fn thread_policy(&self) -> ThreadPolicy {
self.threads
}
pub fn jit_policy(&self) -> JitPolicy {
self.jit
}
pub fn partitioning(&self) -> Partitioning {
self.partitioning
}
pub fn memory_state(&self) -> &MemoryState {
&self.memory_state
}
pub fn host_memory(&self) -> &MemoryPool {
&self.host_memory
}
pub fn device_memory(&self) -> Option<&MemoryPool> {
self.device_memory.as_ref()
}
pub fn memory_report(&self) -> MemoryReport {
self.memory_state.report()
}
pub fn memory_pool_reports(&self) -> Vec<MemoryPoolReport> {
std::iter::once(self.host_memory.report())
.chain(self.device_memory.as_ref().map(MemoryPool::report))
.collect()
}
pub fn memory_decisions(&self) -> Vec<MemoryDecision> {
self.memory_decisions
.lock()
.unwrap_or_else(|error| error.into_inner())
.clone()
}
pub fn record_memory_decision(&self, decision: MemoryDecision) {
self.memory_decisions
.lock()
.unwrap_or_else(|error| error.into_inner())
.push(decision);
}
pub fn rank(&self) -> usize {
#[cfg(feature = "mpi")]
if let Some(communicator) = &self.communicator {
return communicator.rank() as usize;
}
0
}
pub fn nranks(&self) -> usize {
#[cfg(feature = "mpi")]
if let Some(communicator) = &self.communicator {
return communicator.size() as usize;
}
1
}
pub fn is_distributed(&self) -> bool {
self.nranks() > 1
}
#[allow(unused_mut)]
pub(crate) fn read_plan(&self, mut plan: ReadPlan) -> ReadPlan {
#[cfg(feature = "mpi")]
if let Some(communicator) = &self.communicator {
plan.distribution = laddu_data::io::Distribution::from_world(communicator.as_ref())
.with_partitioning(self.partitioning);
}
plan
}
pub(crate) fn sum_f64(&self, local: f64) -> f64 {
#[cfg(feature = "mpi")]
if let Some(communicator) = &self.communicator {
let mut global = 0.0;
communicator.all_reduce_into(&local, &mut global, SystemOperation::sum());
return global;
}
local
}
pub(crate) fn sum_usize(&self, local: usize) -> usize {
#[cfg(feature = "mpi")]
if let Some(communicator) = &self.communicator {
let local = local as u64;
let mut global = 0_u64;
communicator.all_reduce_into(&local, &mut global, SystemOperation::sum());
return global as usize;
}
local
}
pub(crate) fn sum_slice(&self, local: &[f64]) -> Vec<f64> {
#[cfg(feature = "mpi")]
if let Some(communicator) = &self.communicator {
let mut global = vec![0.0; local.len()];
communicator.all_reduce_into(local, &mut global, SystemOperation::sum());
return global;
}
local.to_vec()
}
pub(crate) fn all_succeeded(&self, local_success: bool) -> bool {
self.sum_usize(usize::from(local_success)) == self.nranks()
}
pub(crate) fn is_parallel(&self) -> bool {
self.threads != ThreadPolicy::Serial
}
pub(crate) fn install<R: Send>(&self, operation: impl FnOnce() -> R + Send) -> R {
match &self.pool {
Some(pool) => pool.install(operation),
None => operation(),
}
}
}
#[cfg(feature = "mpi")]
fn shared_mpi_budget(budget: MemoryBudget, local_processes: u64) -> MemoryBudget {
let divisor = local_processes.max(1);
match budget {
MemoryBudget::Auto => MemoryBudget::PercentAvailable(0.80 / divisor as f64),
MemoryBudget::Bytes(bytes) => MemoryBudget::Bytes((bytes / divisor).max(1)),
MemoryBudget::PercentTotal(fraction) => {
MemoryBudget::PercentTotal(fraction / divisor as f64)
}
MemoryBudget::PercentAvailable(fraction) => {
MemoryBudget::PercentAvailable(fraction / divisor as f64)
}
}
}
#[cfg(feature = "mpi")]
fn mpi_local_process_count(world_size: i32) -> u64 {
const VARIABLES: [&str; 4] = [
"OMPI_COMM_WORLD_LOCAL_SIZE",
"MPI_LOCALNRANKS",
"MV2_COMM_WORLD_LOCAL_SIZE",
"SLURM_NTASKS_PER_NODE",
];
VARIABLES
.iter()
.filter_map(|name| std::env::var(name).ok())
.find_map(|value| {
value
.split(|character: char| !character.is_ascii_digit())
.find(|part| !part.is_empty())
.and_then(|part| part.parse::<u64>().ok())
.filter(|count| *count > 0)
})
.unwrap_or_else(|| u64::try_from(world_size).unwrap_or(1).max(1))
}
#[cfg(test)]
mod tests {
use super::*;
#[cfg(not(feature = "wgpu"))]
use crate::RuntimeError;
use crate::execution::GpuBackend;
#[test]
fn execution_options_roundtrip_through_json() {
let options = ExecutionOptions {
device: Device::Gpu(GpuOptions {
backend: GpuBackend::Wgpu,
device: GpuDeviceSelector::PciBusId("0000:01:00.0".into()),
}),
precision: Precision::F64,
autodiff: AutodiffMode::Reverse,
normalization: NormalizationMode::Verify,
partitioning: Partitioning::FileGroups,
memory: MemoryPlan::host_device(
MemoryBudget::PercentAvailable(0.5),
MemoryBudget::Bytes(1 << 30),
),
};
let json = serde_json::to_string(&options).unwrap();
assert_eq!(
serde_json::from_str::<ExecutionOptions>(&json).unwrap(),
options
);
}
#[test]
fn execution_selects_nested_cpu_options() {
let serial = Execution::local(ExecutionOptions {
device: Device::Cpu(CpuOptions {
threads: ThreadPolicy::Serial,
jit: JitPolicy::Disabled,
}),
..ExecutionOptions::default()
})
.unwrap();
assert!(!serial.is_parallel());
assert_eq!(serial.jit_policy(), JitPolicy::Disabled);
assert_eq!(serial.precision(), Precision::F64);
let fixed = Execution::local(ExecutionOptions {
device: Device::Cpu(CpuOptions {
threads: ThreadPolicy::Fixed(2),
..CpuOptions::default()
}),
..ExecutionOptions::default()
})
.unwrap();
assert_eq!(fixed.install(rayon::current_num_threads), 2);
}
#[test]
fn unavailable_execution_modes_return_capability_errors() {
#[cfg(not(feature = "wgpu"))]
assert!(matches!(
Execution::local(ExecutionOptions {
device: Device::Gpu(GpuOptions {
backend: GpuBackend::Wgpu,
..GpuOptions::default()
}),
..ExecutionOptions::default()
}),
Err(RuntimeError::Execution(ExecutionError::GpuUnavailable(
GpuBackend::Wgpu
)))
));
#[cfg(feature = "wgpu")]
assert!(
Execution::local(ExecutionOptions {
device: Device::Gpu(GpuOptions {
backend: GpuBackend::Wgpu,
..GpuOptions::default()
}),
..ExecutionOptions::default()
})
.is_ok()
);
let f32 = Execution::local(ExecutionOptions {
device: Device::Cpu(CpuOptions::default()),
precision: Precision::F32,
..ExecutionOptions::default()
})
.unwrap();
assert_eq!(f32.precision(), Precision::F32);
let reverse = Execution::local(ExecutionOptions {
autodiff: AutodiffMode::Reverse,
..ExecutionOptions::default()
})
.unwrap();
assert_eq!(reverse.autodiff_mode(), AutodiffMode::Reverse);
let reverse_f32 = Execution::local(ExecutionOptions {
precision: Precision::F32,
autodiff: AutodiffMode::Reverse,
..ExecutionOptions::default()
})
.unwrap();
assert_eq!(reverse_f32.precision(), Precision::F32);
assert_eq!(reverse_f32.autodiff_mode(), AutodiffMode::Reverse);
}
}