Skip to main content

sim_lib_compute_wgpu/
probe.rs

1//! Portable GPU adapter discovery and raw probe evidence.
2
3use std::sync::mpsc;
4
5use wgpu::{
6    Adapter, Backends, BufferDescriptor, BufferUsages, DeviceDescriptor, ExperimentalFeatures,
7    Features, Instance, Limits, MapMode, MemoryHints, PollType,
8};
9
10/// Requested limits and optional features passed to `wgpu`.
11#[derive(Clone, Debug, PartialEq, Eq)]
12pub struct RequestedWgpuProfile {
13    /// Requested device limits.
14    pub limits: WgpuLimitEvidence,
15    /// Whether timestamp queries were requested.
16    pub timestamp_query: bool,
17    /// Whether shader f16 was requested.
18    pub shader_f16: bool,
19}
20
21impl RequestedWgpuProfile {
22    fn from_parts(limits: Limits, features: Features) -> Self {
23        Self {
24            limits: WgpuLimitEvidence::from_limits(&limits),
25            timestamp_query: features.contains(Features::TIMESTAMP_QUERY),
26            shader_f16: features.contains(Features::SHADER_F16),
27        }
28    }
29}
30
31/// Limits recorded from either the requested or granted device contract.
32#[derive(Clone, Debug, PartialEq, Eq)]
33pub struct WgpuLimitEvidence {
34    /// Maximum buffer size.
35    pub max_buffer_size: u64,
36    /// Maximum storage buffer binding size.
37    pub max_storage_buffer_binding_size: u64,
38    /// Maximum uniform buffer binding size.
39    pub max_uniform_buffer_binding_size: u64,
40    /// Minimum storage buffer offset alignment.
41    pub min_storage_buffer_offset_alignment: u32,
42    /// Minimum uniform buffer offset alignment.
43    pub min_uniform_buffer_offset_alignment: u32,
44    /// Maximum compute workgroups per dimension.
45    pub max_compute_workgroups_per_dimension: u32,
46    /// Maximum compute invocations per workgroup.
47    pub max_compute_invocations_per_workgroup: u32,
48    /// Maximum compute workgroup size x.
49    pub max_compute_workgroup_size_x: u32,
50    /// Maximum compute workgroup size y.
51    pub max_compute_workgroup_size_y: u32,
52    /// Maximum compute workgroup size z.
53    pub max_compute_workgroup_size_z: u32,
54}
55
56impl WgpuLimitEvidence {
57    fn from_limits(limits: &Limits) -> Self {
58        Self {
59            max_buffer_size: limits.max_buffer_size,
60            max_storage_buffer_binding_size: limits.max_storage_buffer_binding_size,
61            max_uniform_buffer_binding_size: limits.max_uniform_buffer_binding_size,
62            min_storage_buffer_offset_alignment: limits.min_storage_buffer_offset_alignment,
63            min_uniform_buffer_offset_alignment: limits.min_uniform_buffer_offset_alignment,
64            max_compute_workgroups_per_dimension: limits.max_compute_workgroups_per_dimension,
65            max_compute_invocations_per_workgroup: limits.max_compute_invocations_per_workgroup,
66            max_compute_workgroup_size_x: limits.max_compute_workgroup_size_x,
67            max_compute_workgroup_size_y: limits.max_compute_workgroup_size_y,
68            max_compute_workgroup_size_z: limits.max_compute_workgroup_size_z,
69        }
70    }
71}
72
73/// Feature evidence recorded from a granted device.
74#[derive(Clone, Debug, PartialEq, Eq)]
75pub struct WgpuCapabilityEvidence {
76    /// Whether timestamp queries are granted.
77    pub timestamp_query: bool,
78    /// Whether shader f16 is granted.
79    pub shader_f16: bool,
80    /// Whether primary buffers may be mapped.
81    pub mappable_primary_buffers: bool,
82}
83
84impl WgpuCapabilityEvidence {
85    fn from_features(features: Features) -> Self {
86        Self {
87            timestamp_query: features.contains(Features::TIMESTAMP_QUERY),
88            shader_f16: features.contains(Features::SHADER_F16),
89            mappable_primary_buffers: features.contains(Features::MAPPABLE_PRIMARY_BUFFERS),
90        }
91    }
92}
93
94/// Adapter identity and requested/granted capability evidence.
95#[derive(Clone, Debug, PartialEq, Eq)]
96pub struct WgpuAdapterEvidence {
97    /// Deterministic ordinal assigned after sorting adapters.
98    pub ordinal: usize,
99    /// Diagnostic adapter name from `wgpu`.
100    pub name: String,
101    /// Diagnostic backend label from `wgpu`.
102    pub backend: String,
103    /// Diagnostic adapter type from `wgpu`.
104    pub adapter_type: String,
105    /// Diagnostic vendor id.
106    pub vendor: u32,
107    /// Diagnostic device id.
108    pub device: u32,
109    /// Requested profile.
110    pub requested: RequestedWgpuProfile,
111    /// Granted device limits.
112    pub granted_limits: WgpuLimitEvidence,
113    /// Granted features.
114    pub granted_features: WgpuCapabilityEvidence,
115}
116
117impl WgpuAdapterEvidence {
118    /// Sort key that keeps enumeration deterministic without treating identity
119    /// as product logic.
120    pub fn sort_key(&self) -> (&str, &str, &str, u32, u32) {
121        (
122            self.backend.as_str(),
123            self.adapter_type.as_str(),
124            self.name.as_str(),
125            self.vendor,
126            self.device,
127        )
128    }
129}
130
131/// Transfer and mapping probe evidence.
132#[derive(Clone, Debug, PartialEq, Eq)]
133pub struct TransferEvidence {
134    /// Bytes written and read back.
135    pub bytes: u64,
136    /// Whether queue write plus copy completed.
137    pub transfer_ok: bool,
138    /// Whether map-read completed and matched the payload.
139    pub mapping_ok: bool,
140}
141
142/// One bounded allocation attempt.
143#[derive(Clone, Debug, PartialEq, Eq)]
144pub struct AllocationAttempt {
145    /// Attempted byte size.
146    pub bytes: u64,
147    /// Whether creating the buffer succeeded.
148    pub success: bool,
149}
150
151/// Probe evidence required before a site is exported.
152#[derive(Clone, Debug, PartialEq, Eq)]
153pub struct ProbeEvidence {
154    /// Transfer and mapping evidence.
155    pub transfer: TransferEvidence,
156    /// Bounded allocation attempts.
157    pub allocation_attempts: Vec<AllocationAttempt>,
158}
159
160impl ProbeEvidence {
161    /// Returns true when all required probes succeeded.
162    pub fn successful(&self) -> bool {
163        self.transfer.transfer_ok
164            && self.transfer.mapping_ok
165            && self
166                .allocation_attempts
167                .iter()
168                .any(|attempt| attempt.success && attempt.bytes > 0)
169    }
170}
171
172/// One successful adapter probe.
173#[derive(Clone, Debug, PartialEq, Eq)]
174pub struct WgpuAdapterProbe {
175    /// Adapter and capability evidence.
176    pub adapter: WgpuAdapterEvidence,
177    /// Raw probe evidence.
178    pub probe: ProbeEvidence,
179}
180
181/// Complete discovery result.
182#[derive(Clone, Debug, Default, PartialEq, Eq)]
183pub struct WgpuDiscovery {
184    /// Successful adapter probes, in deterministic order.
185    pub adapters: Vec<WgpuAdapterProbe>,
186    /// Diagnostic errors from adapters that did not become sites.
187    pub diagnostics: Vec<String>,
188}
189
190impl WgpuDiscovery {
191    /// Builds a discovery result and drops unsuccessful adapters from the site
192    /// list while preserving their diagnostics.
193    pub fn from_probes(probes: Vec<WgpuAdapterProbe>, mut diagnostics: Vec<String>) -> Self {
194        let mut adapters = Vec::new();
195        for probe in probes {
196            if probe.probe.successful() {
197                adapters.push(probe);
198            } else {
199                diagnostics.push(format!(
200                    "wgpu adapter {} did not pass required probes",
201                    probe.adapter.name
202                ));
203            }
204        }
205        adapters.sort_by(|left, right| left.adapter.sort_key().cmp(&right.adapter.sort_key()));
206        for (ordinal, probe) in adapters.iter_mut().enumerate() {
207            probe.adapter.ordinal = ordinal;
208        }
209        Self {
210            adapters,
211            diagnostics,
212        }
213    }
214}
215
216/// Bounded probe policy.
217#[derive(Clone, Debug, PartialEq, Eq)]
218pub struct ProbePolicy {
219    /// Backends to enumerate.
220    pub backends: Backends,
221    /// Bytes used by the transfer and map probe.
222    pub transfer_bytes: u64,
223    /// Largest allocation attempt, capped again by granted limits.
224    pub max_allocation_probe_bytes: u64,
225}
226
227impl Default for ProbePolicy {
228    fn default() -> Self {
229        Self {
230            backends: Backends::all(),
231            transfer_bytes: 16,
232            max_allocation_probe_bytes: 16 * 1024 * 1024,
233        }
234    }
235}
236
237/// Discovery failure for infrastructure-level probe setup.
238#[derive(Clone, Debug, PartialEq, Eq)]
239pub struct WgpuDiscoveryError {
240    message: String,
241}
242
243impl WgpuDiscoveryError {
244    fn new(message: impl Into<String>) -> Self {
245        Self {
246            message: message.into(),
247        }
248    }
249}
250
251impl std::fmt::Display for WgpuDiscoveryError {
252    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
253        formatter.write_str(&self.message)
254    }
255}
256
257impl std::error::Error for WgpuDiscoveryError {}
258
259/// Enumerates adapters and returns only probe-backed site candidates.
260pub fn discover_wgpu_adapters(policy: &ProbePolicy) -> Result<WgpuDiscovery, WgpuDiscoveryError> {
261    let instance = Instance::default();
262    let adapters = pollster::block_on(instance.enumerate_adapters(policy.backends));
263    let mut probes = Vec::new();
264    let mut diagnostics = Vec::new();
265
266    for adapter in adapters {
267        match probe_adapter(adapter, policy) {
268            Ok(probe) => probes.push(probe),
269            Err(error) => diagnostics.push(error.to_string()),
270        }
271    }
272
273    Ok(WgpuDiscovery::from_probes(probes, diagnostics))
274}
275
276fn probe_adapter(
277    adapter: Adapter,
278    policy: &ProbePolicy,
279) -> Result<WgpuAdapterProbe, WgpuDiscoveryError> {
280    let info = adapter.get_info();
281    let supported_features = adapter.features();
282    let required_features = supported_features & (Features::TIMESTAMP_QUERY | Features::SHADER_F16);
283    let required_limits = Limits::downlevel_defaults().using_resolution(adapter.limits());
284    let requested = RequestedWgpuProfile::from_parts(required_limits.clone(), required_features);
285    let descriptor = DeviceDescriptor {
286        label: Some("sim-compute-wgpu-probe"),
287        required_features,
288        required_limits,
289        experimental_features: ExperimentalFeatures::disabled(),
290        memory_hints: MemoryHints::Performance,
291        trace: Default::default(),
292    };
293    let (device, queue) = pollster::block_on(adapter.request_device(&descriptor))
294        .map_err(|err| WgpuDiscoveryError::new(format!("wgpu request_device failed: {err}")))?;
295
296    let transfer = probe_transfer(&device, &queue, policy.transfer_bytes)?;
297    let allocation_attempts = probe_allocations(
298        &device,
299        device
300            .limits()
301            .max_buffer_size
302            .min(policy.max_allocation_probe_bytes),
303    );
304
305    Ok(WgpuAdapterProbe {
306        adapter: WgpuAdapterEvidence {
307            ordinal: 0,
308            name: info.name,
309            backend: format!("{:?}", info.backend),
310            adapter_type: format!("{:?}", info.device_type),
311            vendor: info.vendor,
312            device: info.device,
313            requested,
314            granted_limits: WgpuLimitEvidence::from_limits(&device.limits()),
315            granted_features: WgpuCapabilityEvidence::from_features(device.features()),
316        },
317        probe: ProbeEvidence {
318            transfer,
319            allocation_attempts,
320        },
321    })
322}
323
324fn probe_transfer(
325    device: &wgpu::Device,
326    queue: &wgpu::Queue,
327    bytes: u64,
328) -> Result<TransferEvidence, WgpuDiscoveryError> {
329    let bytes = bytes.max(4).next_multiple_of(4);
330    let payload = (0..bytes).map(|idx| (idx % 251) as u8).collect::<Vec<_>>();
331    let source = device.create_buffer(&BufferDescriptor {
332        label: Some("sim-compute-wgpu-transfer-source"),
333        size: bytes,
334        usage: BufferUsages::COPY_SRC | BufferUsages::COPY_DST,
335        mapped_at_creation: false,
336    });
337    let readback = device.create_buffer(&BufferDescriptor {
338        label: Some("sim-compute-wgpu-transfer-readback"),
339        size: bytes,
340        usage: BufferUsages::COPY_DST | BufferUsages::MAP_READ,
341        mapped_at_creation: false,
342    });
343    queue.write_buffer(&source, 0, &payload);
344    let mut encoder = device.create_command_encoder(&wgpu::CommandEncoderDescriptor {
345        label: Some("sim-compute-wgpu-transfer-encoder"),
346    });
347    encoder.copy_buffer_to_buffer(&source, 0, &readback, 0, bytes);
348    queue.submit([encoder.finish()]);
349
350    let (sender, receiver) = mpsc::channel();
351    readback.slice(..).map_async(MapMode::Read, move |result| {
352        let _ = sender.send(result);
353    });
354    device
355        .poll(PollType::wait_indefinitely())
356        .map_err(|err| WgpuDiscoveryError::new(format!("wgpu poll failed: {err}")))?;
357    receiver
358        .recv()
359        .map_err(|err| WgpuDiscoveryError::new(format!("wgpu map callback failed: {err}")))?
360        .map_err(|err| WgpuDiscoveryError::new(format!("wgpu map failed: {err}")))?;
361
362    let mapped = readback
363        .slice(..)
364        .get_mapped_range()
365        .map_err(|err| WgpuDiscoveryError::new(format!("wgpu mapped range failed: {err}")))?
366        .to_vec();
367    readback.unmap();
368    Ok(TransferEvidence {
369        bytes,
370        transfer_ok: true,
371        mapping_ok: mapped == payload,
372    })
373}
374
375fn probe_allocations(device: &wgpu::Device, ceiling: u64) -> Vec<AllocationAttempt> {
376    [4096, 1024 * 1024, ceiling]
377        .into_iter()
378        .filter(|bytes| *bytes > 0)
379        .map(|bytes| {
380            let buffer = device.create_buffer(&BufferDescriptor {
381                label: Some("sim-compute-wgpu-allocation-probe"),
382                size: bytes,
383                usage: BufferUsages::COPY_DST,
384                mapped_at_creation: false,
385            });
386            drop(buffer);
387            AllocationAttempt {
388                bytes,
389                success: true,
390            }
391        })
392        .collect()
393}