use std::sync::mpsc;
use sim_lib_compute_auto::{ComputeDeviceIdentity, ComputeEvidenceKind, ComputePhysicalEvidence};
use wgpu::{
Adapter, Backends, BufferDescriptor, BufferUsages, DeviceDescriptor, ExperimentalFeatures,
Features, Instance, Limits, MapMode, MemoryHints, PollType,
};
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct RequestedWgpuProfile {
pub limits: WgpuLimitEvidence,
pub timestamp_query: bool,
pub shader_f16: bool,
}
impl RequestedWgpuProfile {
fn from_parts(limits: Limits, features: Features) -> Self {
Self {
limits: WgpuLimitEvidence::from_limits(&limits),
timestamp_query: features.contains(Features::TIMESTAMP_QUERY),
shader_f16: features.contains(Features::SHADER_F16),
}
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct WgpuLimitEvidence {
pub max_buffer_size: u64,
pub max_storage_buffer_binding_size: u64,
pub max_uniform_buffer_binding_size: u64,
pub min_storage_buffer_offset_alignment: u32,
pub min_uniform_buffer_offset_alignment: u32,
pub max_compute_workgroups_per_dimension: u32,
pub max_compute_invocations_per_workgroup: u32,
pub max_compute_workgroup_size_x: u32,
pub max_compute_workgroup_size_y: u32,
pub max_compute_workgroup_size_z: u32,
}
impl WgpuLimitEvidence {
fn from_limits(limits: &Limits) -> Self {
Self {
max_buffer_size: limits.max_buffer_size,
max_storage_buffer_binding_size: limits.max_storage_buffer_binding_size,
max_uniform_buffer_binding_size: limits.max_uniform_buffer_binding_size,
min_storage_buffer_offset_alignment: limits.min_storage_buffer_offset_alignment,
min_uniform_buffer_offset_alignment: limits.min_uniform_buffer_offset_alignment,
max_compute_workgroups_per_dimension: limits.max_compute_workgroups_per_dimension,
max_compute_invocations_per_workgroup: limits.max_compute_invocations_per_workgroup,
max_compute_workgroup_size_x: limits.max_compute_workgroup_size_x,
max_compute_workgroup_size_y: limits.max_compute_workgroup_size_y,
max_compute_workgroup_size_z: limits.max_compute_workgroup_size_z,
}
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct WgpuCapabilityEvidence {
pub timestamp_query: bool,
pub shader_f16: bool,
pub mappable_primary_buffers: bool,
}
impl WgpuCapabilityEvidence {
fn from_features(features: Features) -> Self {
Self {
timestamp_query: features.contains(Features::TIMESTAMP_QUERY),
shader_f16: features.contains(Features::SHADER_F16),
mappable_primary_buffers: features.contains(Features::MAPPABLE_PRIMARY_BUFFERS),
}
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct WgpuAdapterEvidence {
pub ordinal: usize,
pub name: String,
pub backend: String,
pub adapter_type: String,
pub vendor: u32,
pub device: u32,
pub requested: RequestedWgpuProfile,
pub granted_limits: WgpuLimitEvidence,
pub granted_features: WgpuCapabilityEvidence,
}
impl WgpuAdapterEvidence {
pub fn sort_key(&self) -> (&str, &str, &str, u32, u32) {
(
self.backend.as_str(),
self.adapter_type.as_str(),
self.name.as_str(),
self.vendor,
self.device,
)
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct TransferEvidence {
pub bytes: u64,
pub transfer_ok: bool,
pub mapping_ok: bool,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct AllocationAttempt {
pub bytes: u64,
pub success: bool,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ProbeEvidence {
pub transfer: TransferEvidence,
pub allocation_attempts: Vec<AllocationAttempt>,
}
impl ProbeEvidence {
pub fn successful(&self) -> bool {
self.transfer.transfer_ok
&& self.transfer.mapping_ok
&& self
.allocation_attempts
.iter()
.any(|attempt| attempt.success && attempt.bytes > 0)
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct WgpuAdapterProbe {
pub evidence_kind: ComputeEvidenceKind,
pub claimed_identity: Option<ComputeDeviceIdentity>,
pub observed_identity: Option<ComputeDeviceIdentity>,
pub adapter: WgpuAdapterEvidence,
pub probe: ProbeEvidence,
}
pub(crate) struct WgpuAdapterRuntime {
pub(crate) probe: WgpuAdapterProbe,
pub(crate) device: wgpu::Device,
pub(crate) queue: wgpu::Queue,
}
impl ComputePhysicalEvidence for WgpuAdapterProbe {
fn evidence_kind(&self) -> ComputeEvidenceKind {
self.evidence_kind
}
fn claimed_identity(&self) -> Option<&ComputeDeviceIdentity> {
self.claimed_identity.as_ref()
}
fn observed_identity(&self) -> Option<&ComputeDeviceIdentity> {
self.observed_identity.as_ref()
}
}
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct WgpuDiscovery {
pub adapters: Vec<WgpuAdapterProbe>,
pub diagnostics: Vec<String>,
}
impl WgpuDiscovery {
pub fn from_probes(probes: Vec<WgpuAdapterProbe>, mut diagnostics: Vec<String>) -> Self {
let mut adapters = Vec::new();
for probe in probes {
if probe.probe.successful() {
adapters.push(probe);
} else {
diagnostics.push(format!(
"wgpu adapter {} did not pass required probes",
probe.adapter.name
));
}
}
adapters.sort_by(|left, right| left.adapter.sort_key().cmp(&right.adapter.sort_key()));
for (ordinal, probe) in adapters.iter_mut().enumerate() {
probe.adapter.ordinal = ordinal;
}
Self {
adapters,
diagnostics,
}
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ProbePolicy {
pub backends: Backends,
pub transfer_bytes: u64,
pub max_allocation_probe_bytes: u64,
}
impl Default for ProbePolicy {
fn default() -> Self {
Self {
backends: Backends::all(),
transfer_bytes: 16,
max_allocation_probe_bytes: 16 * 1024 * 1024,
}
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct WgpuDiscoveryError {
message: String,
}
impl WgpuDiscoveryError {
fn new(message: impl Into<String>) -> Self {
Self {
message: message.into(),
}
}
}
impl std::fmt::Display for WgpuDiscoveryError {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str(&self.message)
}
}
impl std::error::Error for WgpuDiscoveryError {}
pub fn discover_wgpu_adapters(policy: &ProbePolicy) -> Result<WgpuDiscovery, WgpuDiscoveryError> {
let (runtimes, diagnostics) = discover_wgpu_adapter_runtimes_with_diagnostics(policy)?;
Ok(WgpuDiscovery::from_probes(
runtimes.into_iter().map(|runtime| runtime.probe).collect(),
diagnostics,
))
}
pub(crate) fn discover_wgpu_adapter_runtimes(
policy: &ProbePolicy,
) -> Result<Vec<WgpuAdapterRuntime>, WgpuDiscoveryError> {
discover_wgpu_adapter_runtimes_with_diagnostics(policy).map(|(runtimes, _)| runtimes)
}
fn discover_wgpu_adapter_runtimes_with_diagnostics(
policy: &ProbePolicy,
) -> Result<(Vec<WgpuAdapterRuntime>, Vec<String>), WgpuDiscoveryError> {
let instance = Instance::default();
let adapters = pollster::block_on(instance.enumerate_adapters(policy.backends));
let mut runtimes = Vec::new();
let mut diagnostics = Vec::new();
for adapter in adapters {
match probe_adapter(adapter, policy) {
Ok(runtime) => runtimes.push(runtime),
Err(error) => diagnostics.push(error.to_string()),
}
}
runtimes.retain(|runtime| {
if runtime.probe.probe.successful() {
true
} else {
diagnostics.push(format!(
"wgpu adapter {} did not pass required probes",
runtime.probe.adapter.name
));
false
}
});
runtimes.sort_by(|left, right| {
left.probe
.adapter
.sort_key()
.cmp(&right.probe.adapter.sort_key())
});
for (ordinal, runtime) in runtimes.iter_mut().enumerate() {
runtime.probe.adapter.ordinal = ordinal;
}
Ok((runtimes, diagnostics))
}
fn probe_adapter(
adapter: Adapter,
policy: &ProbePolicy,
) -> Result<WgpuAdapterRuntime, WgpuDiscoveryError> {
let info = adapter.get_info();
let backend = format!("{:?}", info.backend);
let identity = ComputeDeviceIdentity::new(info.name.clone(), "wgpu", backend.clone());
let supported_features = adapter.features();
let required_features = supported_features & (Features::TIMESTAMP_QUERY | Features::SHADER_F16);
let required_limits = Limits::downlevel_defaults().using_resolution(adapter.limits());
let requested = RequestedWgpuProfile::from_parts(required_limits.clone(), required_features);
let descriptor = DeviceDescriptor {
label: Some("sim-compute-wgpu-probe"),
required_features,
required_limits,
experimental_features: ExperimentalFeatures::disabled(),
memory_hints: MemoryHints::Performance,
trace: Default::default(),
};
let (device, queue) = pollster::block_on(adapter.request_device(&descriptor))
.map_err(|err| WgpuDiscoveryError::new(format!("wgpu request_device failed: {err}")))?;
let transfer = probe_transfer(&device, &queue, policy.transfer_bytes)?;
let allocation_attempts = probe_allocations(
&device,
device
.limits()
.max_buffer_size
.min(policy.max_allocation_probe_bytes),
);
Ok(WgpuAdapterRuntime {
probe: WgpuAdapterProbe {
evidence_kind: ComputeEvidenceKind::PhysicalDevice,
claimed_identity: Some(identity.clone()),
observed_identity: Some(identity),
adapter: WgpuAdapterEvidence {
ordinal: 0,
name: info.name,
backend,
adapter_type: format!("{:?}", info.device_type),
vendor: info.vendor,
device: info.device,
requested,
granted_limits: WgpuLimitEvidence::from_limits(&device.limits()),
granted_features: WgpuCapabilityEvidence::from_features(device.features()),
},
probe: ProbeEvidence {
transfer,
allocation_attempts,
},
},
device,
queue,
})
}
fn probe_transfer(
device: &wgpu::Device,
queue: &wgpu::Queue,
bytes: u64,
) -> Result<TransferEvidence, WgpuDiscoveryError> {
let bytes = bytes.max(4).next_multiple_of(4);
let payload = (0..bytes).map(|idx| (idx % 251) as u8).collect::<Vec<_>>();
let source = device.create_buffer(&BufferDescriptor {
label: Some("sim-compute-wgpu-transfer-source"),
size: bytes,
usage: BufferUsages::COPY_SRC | BufferUsages::COPY_DST,
mapped_at_creation: false,
});
let readback = device.create_buffer(&BufferDescriptor {
label: Some("sim-compute-wgpu-transfer-readback"),
size: bytes,
usage: BufferUsages::COPY_DST | BufferUsages::MAP_READ,
mapped_at_creation: false,
});
queue.write_buffer(&source, 0, &payload);
let mut encoder = device.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("sim-compute-wgpu-transfer-encoder"),
});
encoder.copy_buffer_to_buffer(&source, 0, &readback, 0, bytes);
queue.submit([encoder.finish()]);
let (sender, receiver) = mpsc::channel();
readback.slice(..).map_async(MapMode::Read, move |result| {
let _ = sender.send(result);
});
device
.poll(PollType::wait_indefinitely())
.map_err(|err| WgpuDiscoveryError::new(format!("wgpu poll failed: {err}")))?;
receiver
.recv()
.map_err(|err| WgpuDiscoveryError::new(format!("wgpu map callback failed: {err}")))?
.map_err(|err| WgpuDiscoveryError::new(format!("wgpu map failed: {err}")))?;
let mapped = readback
.slice(..)
.get_mapped_range()
.map_err(|err| WgpuDiscoveryError::new(format!("wgpu mapped range failed: {err}")))?
.to_vec();
readback.unmap();
Ok(TransferEvidence {
bytes,
transfer_ok: true,
mapping_ok: mapped == payload,
})
}
fn probe_allocations(device: &wgpu::Device, ceiling: u64) -> Vec<AllocationAttempt> {
[4096, 1024 * 1024, ceiling]
.into_iter()
.filter(|bytes| *bytes > 0)
.map(|bytes| {
let buffer = device.create_buffer(&BufferDescriptor {
label: Some("sim-compute-wgpu-allocation-probe"),
size: bytes,
usage: BufferUsages::COPY_DST,
mapped_at_creation: false,
});
drop(buffer);
AllocationAttempt {
bytes,
success: true,
}
})
.chain(std::iter::once(AllocationAttempt {
bytes: ceiling.saturating_add(1),
success: false,
}))
.collect()
}