1use std::sync::mpsc;
4
5use wgpu::{
6 Adapter, Backends, BufferDescriptor, BufferUsages, DeviceDescriptor, ExperimentalFeatures,
7 Features, Instance, Limits, MapMode, MemoryHints, PollType,
8};
9
10#[derive(Clone, Debug, PartialEq, Eq)]
12pub struct RequestedWgpuProfile {
13 pub limits: WgpuLimitEvidence,
15 pub timestamp_query: bool,
17 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#[derive(Clone, Debug, PartialEq, Eq)]
33pub struct WgpuLimitEvidence {
34 pub max_buffer_size: u64,
36 pub max_storage_buffer_binding_size: u64,
38 pub max_uniform_buffer_binding_size: u64,
40 pub min_storage_buffer_offset_alignment: u32,
42 pub min_uniform_buffer_offset_alignment: u32,
44 pub max_compute_workgroups_per_dimension: u32,
46 pub max_compute_invocations_per_workgroup: u32,
48 pub max_compute_workgroup_size_x: u32,
50 pub max_compute_workgroup_size_y: u32,
52 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#[derive(Clone, Debug, PartialEq, Eq)]
75pub struct WgpuCapabilityEvidence {
76 pub timestamp_query: bool,
78 pub shader_f16: bool,
80 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#[derive(Clone, Debug, PartialEq, Eq)]
96pub struct WgpuAdapterEvidence {
97 pub ordinal: usize,
99 pub name: String,
101 pub backend: String,
103 pub adapter_type: String,
105 pub vendor: u32,
107 pub device: u32,
109 pub requested: RequestedWgpuProfile,
111 pub granted_limits: WgpuLimitEvidence,
113 pub granted_features: WgpuCapabilityEvidence,
115}
116
117impl WgpuAdapterEvidence {
118 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#[derive(Clone, Debug, PartialEq, Eq)]
133pub struct TransferEvidence {
134 pub bytes: u64,
136 pub transfer_ok: bool,
138 pub mapping_ok: bool,
140}
141
142#[derive(Clone, Debug, PartialEq, Eq)]
144pub struct AllocationAttempt {
145 pub bytes: u64,
147 pub success: bool,
149}
150
151#[derive(Clone, Debug, PartialEq, Eq)]
153pub struct ProbeEvidence {
154 pub transfer: TransferEvidence,
156 pub allocation_attempts: Vec<AllocationAttempt>,
158}
159
160impl ProbeEvidence {
161 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#[derive(Clone, Debug, PartialEq, Eq)]
174pub struct WgpuAdapterProbe {
175 pub adapter: WgpuAdapterEvidence,
177 pub probe: ProbeEvidence,
179}
180
181#[derive(Clone, Debug, Default, PartialEq, Eq)]
183pub struct WgpuDiscovery {
184 pub adapters: Vec<WgpuAdapterProbe>,
186 pub diagnostics: Vec<String>,
188}
189
190impl WgpuDiscovery {
191 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#[derive(Clone, Debug, PartialEq, Eq)]
218pub struct ProbePolicy {
219 pub backends: Backends,
221 pub transfer_bytes: u64,
223 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#[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
259pub 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}