1use std::sync::mpsc;
4
5use sim_lib_compute_auto::{ComputeDeviceIdentity, ComputeEvidenceKind, ComputePhysicalEvidence};
6use wgpu::{
7 Adapter, Backends, BufferDescriptor, BufferUsages, DeviceDescriptor, ExperimentalFeatures,
8 Features, Instance, Limits, MapMode, MemoryHints, PollType,
9};
10
11#[derive(Clone, Debug, PartialEq, Eq)]
13pub struct RequestedWgpuProfile {
14 pub limits: WgpuLimitEvidence,
16 pub timestamp_query: bool,
18 pub shader_f16: bool,
20}
21
22impl RequestedWgpuProfile {
23 fn from_parts(limits: Limits, features: Features) -> Self {
24 Self {
25 limits: WgpuLimitEvidence::from_limits(&limits),
26 timestamp_query: features.contains(Features::TIMESTAMP_QUERY),
27 shader_f16: features.contains(Features::SHADER_F16),
28 }
29 }
30}
31
32#[derive(Clone, Debug, PartialEq, Eq)]
34pub struct WgpuLimitEvidence {
35 pub max_buffer_size: u64,
37 pub max_storage_buffer_binding_size: u64,
39 pub max_uniform_buffer_binding_size: u64,
41 pub min_storage_buffer_offset_alignment: u32,
43 pub min_uniform_buffer_offset_alignment: u32,
45 pub max_compute_workgroups_per_dimension: u32,
47 pub max_compute_invocations_per_workgroup: u32,
49 pub max_compute_workgroup_size_x: u32,
51 pub max_compute_workgroup_size_y: u32,
53 pub max_compute_workgroup_size_z: u32,
55}
56
57impl WgpuLimitEvidence {
58 fn from_limits(limits: &Limits) -> Self {
59 Self {
60 max_buffer_size: limits.max_buffer_size,
61 max_storage_buffer_binding_size: limits.max_storage_buffer_binding_size,
62 max_uniform_buffer_binding_size: limits.max_uniform_buffer_binding_size,
63 min_storage_buffer_offset_alignment: limits.min_storage_buffer_offset_alignment,
64 min_uniform_buffer_offset_alignment: limits.min_uniform_buffer_offset_alignment,
65 max_compute_workgroups_per_dimension: limits.max_compute_workgroups_per_dimension,
66 max_compute_invocations_per_workgroup: limits.max_compute_invocations_per_workgroup,
67 max_compute_workgroup_size_x: limits.max_compute_workgroup_size_x,
68 max_compute_workgroup_size_y: limits.max_compute_workgroup_size_y,
69 max_compute_workgroup_size_z: limits.max_compute_workgroup_size_z,
70 }
71 }
72}
73
74#[derive(Clone, Debug, PartialEq, Eq)]
76pub struct WgpuCapabilityEvidence {
77 pub timestamp_query: bool,
79 pub shader_f16: bool,
81 pub mappable_primary_buffers: bool,
83}
84
85impl WgpuCapabilityEvidence {
86 fn from_features(features: Features) -> Self {
87 Self {
88 timestamp_query: features.contains(Features::TIMESTAMP_QUERY),
89 shader_f16: features.contains(Features::SHADER_F16),
90 mappable_primary_buffers: features.contains(Features::MAPPABLE_PRIMARY_BUFFERS),
91 }
92 }
93}
94
95#[derive(Clone, Debug, PartialEq, Eq)]
97pub struct WgpuAdapterEvidence {
98 pub ordinal: usize,
100 pub name: String,
102 pub backend: String,
104 pub adapter_type: String,
106 pub vendor: u32,
108 pub device: u32,
110 pub requested: RequestedWgpuProfile,
112 pub granted_limits: WgpuLimitEvidence,
114 pub granted_features: WgpuCapabilityEvidence,
116}
117
118impl WgpuAdapterEvidence {
119 pub fn sort_key(&self) -> (&str, &str, &str, u32, u32) {
122 (
123 self.backend.as_str(),
124 self.adapter_type.as_str(),
125 self.name.as_str(),
126 self.vendor,
127 self.device,
128 )
129 }
130}
131
132#[derive(Clone, Debug, PartialEq, Eq)]
134pub struct TransferEvidence {
135 pub bytes: u64,
137 pub transfer_ok: bool,
139 pub mapping_ok: bool,
141}
142
143#[derive(Clone, Debug, PartialEq, Eq)]
145pub struct AllocationAttempt {
146 pub bytes: u64,
148 pub success: bool,
150}
151
152#[derive(Clone, Debug, PartialEq, Eq)]
154pub struct ProbeEvidence {
155 pub transfer: TransferEvidence,
157 pub allocation_attempts: Vec<AllocationAttempt>,
159}
160
161impl ProbeEvidence {
162 pub fn successful(&self) -> bool {
164 self.transfer.transfer_ok
165 && self.transfer.mapping_ok
166 && self
167 .allocation_attempts
168 .iter()
169 .any(|attempt| attempt.success && attempt.bytes > 0)
170 }
171}
172
173#[derive(Clone, Debug, PartialEq, Eq)]
175pub struct WgpuAdapterProbe {
176 pub evidence_kind: ComputeEvidenceKind,
178 pub claimed_identity: Option<ComputeDeviceIdentity>,
180 pub observed_identity: Option<ComputeDeviceIdentity>,
182 pub adapter: WgpuAdapterEvidence,
184 pub probe: ProbeEvidence,
186}
187
188pub(crate) struct WgpuAdapterRuntime {
190 pub(crate) probe: WgpuAdapterProbe,
191 pub(crate) device: wgpu::Device,
192 pub(crate) queue: wgpu::Queue,
193}
194
195impl ComputePhysicalEvidence for WgpuAdapterProbe {
196 fn evidence_kind(&self) -> ComputeEvidenceKind {
197 self.evidence_kind
198 }
199
200 fn claimed_identity(&self) -> Option<&ComputeDeviceIdentity> {
201 self.claimed_identity.as_ref()
202 }
203
204 fn observed_identity(&self) -> Option<&ComputeDeviceIdentity> {
205 self.observed_identity.as_ref()
206 }
207}
208
209#[derive(Clone, Debug, Default, PartialEq, Eq)]
211pub struct WgpuDiscovery {
212 pub adapters: Vec<WgpuAdapterProbe>,
214 pub diagnostics: Vec<String>,
216}
217
218impl WgpuDiscovery {
219 pub fn from_probes(probes: Vec<WgpuAdapterProbe>, mut diagnostics: Vec<String>) -> Self {
222 let mut adapters = Vec::new();
223 for probe in probes {
224 if probe.probe.successful() {
225 adapters.push(probe);
226 } else {
227 diagnostics.push(format!(
228 "wgpu adapter {} did not pass required probes",
229 probe.adapter.name
230 ));
231 }
232 }
233 adapters.sort_by(|left, right| left.adapter.sort_key().cmp(&right.adapter.sort_key()));
234 for (ordinal, probe) in adapters.iter_mut().enumerate() {
235 probe.adapter.ordinal = ordinal;
236 }
237 Self {
238 adapters,
239 diagnostics,
240 }
241 }
242}
243
244#[derive(Clone, Debug, PartialEq, Eq)]
246pub struct ProbePolicy {
247 pub backends: Backends,
249 pub transfer_bytes: u64,
251 pub max_allocation_probe_bytes: u64,
253}
254
255impl Default for ProbePolicy {
256 fn default() -> Self {
257 Self {
258 backends: Backends::all(),
259 transfer_bytes: 16,
260 max_allocation_probe_bytes: 16 * 1024 * 1024,
261 }
262 }
263}
264
265#[derive(Clone, Debug, PartialEq, Eq)]
267pub struct WgpuDiscoveryError {
268 message: String,
269}
270
271impl WgpuDiscoveryError {
272 fn new(message: impl Into<String>) -> Self {
273 Self {
274 message: message.into(),
275 }
276 }
277}
278
279impl std::fmt::Display for WgpuDiscoveryError {
280 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
281 formatter.write_str(&self.message)
282 }
283}
284
285impl std::error::Error for WgpuDiscoveryError {}
286
287pub fn discover_wgpu_adapters(policy: &ProbePolicy) -> Result<WgpuDiscovery, WgpuDiscoveryError> {
289 let (runtimes, diagnostics) = discover_wgpu_adapter_runtimes_with_diagnostics(policy)?;
290 Ok(WgpuDiscovery::from_probes(
291 runtimes.into_iter().map(|runtime| runtime.probe).collect(),
292 diagnostics,
293 ))
294}
295
296pub(crate) fn discover_wgpu_adapter_runtimes(
298 policy: &ProbePolicy,
299) -> Result<Vec<WgpuAdapterRuntime>, WgpuDiscoveryError> {
300 discover_wgpu_adapter_runtimes_with_diagnostics(policy).map(|(runtimes, _)| runtimes)
301}
302
303fn discover_wgpu_adapter_runtimes_with_diagnostics(
304 policy: &ProbePolicy,
305) -> Result<(Vec<WgpuAdapterRuntime>, Vec<String>), WgpuDiscoveryError> {
306 let instance = Instance::default();
307 let adapters = pollster::block_on(instance.enumerate_adapters(policy.backends));
308 let mut runtimes = Vec::new();
309 let mut diagnostics = Vec::new();
310
311 for adapter in adapters {
312 match probe_adapter(adapter, policy) {
313 Ok(runtime) => runtimes.push(runtime),
314 Err(error) => diagnostics.push(error.to_string()),
315 }
316 }
317
318 runtimes.retain(|runtime| {
319 if runtime.probe.probe.successful() {
320 true
321 } else {
322 diagnostics.push(format!(
323 "wgpu adapter {} did not pass required probes",
324 runtime.probe.adapter.name
325 ));
326 false
327 }
328 });
329 runtimes.sort_by(|left, right| {
330 left.probe
331 .adapter
332 .sort_key()
333 .cmp(&right.probe.adapter.sort_key())
334 });
335 for (ordinal, runtime) in runtimes.iter_mut().enumerate() {
336 runtime.probe.adapter.ordinal = ordinal;
337 }
338 Ok((runtimes, diagnostics))
339}
340
341fn probe_adapter(
342 adapter: Adapter,
343 policy: &ProbePolicy,
344) -> Result<WgpuAdapterRuntime, WgpuDiscoveryError> {
345 let info = adapter.get_info();
346 let backend = format!("{:?}", info.backend);
347 let identity = ComputeDeviceIdentity::new(info.name.clone(), "wgpu", backend.clone());
348 let supported_features = adapter.features();
349 let required_features = supported_features & (Features::TIMESTAMP_QUERY | Features::SHADER_F16);
350 let required_limits = Limits::downlevel_defaults().using_resolution(adapter.limits());
351 let requested = RequestedWgpuProfile::from_parts(required_limits.clone(), required_features);
352 let descriptor = DeviceDescriptor {
353 label: Some("sim-compute-wgpu-probe"),
354 required_features,
355 required_limits,
356 experimental_features: ExperimentalFeatures::disabled(),
357 memory_hints: MemoryHints::Performance,
358 trace: Default::default(),
359 };
360 let (device, queue) = pollster::block_on(adapter.request_device(&descriptor))
361 .map_err(|err| WgpuDiscoveryError::new(format!("wgpu request_device failed: {err}")))?;
362
363 let transfer = probe_transfer(&device, &queue, policy.transfer_bytes)?;
364 let allocation_attempts = probe_allocations(
365 &device,
366 device
367 .limits()
368 .max_buffer_size
369 .min(policy.max_allocation_probe_bytes),
370 );
371
372 Ok(WgpuAdapterRuntime {
373 probe: WgpuAdapterProbe {
374 evidence_kind: ComputeEvidenceKind::PhysicalDevice,
375 claimed_identity: Some(identity.clone()),
376 observed_identity: Some(identity),
377 adapter: WgpuAdapterEvidence {
378 ordinal: 0,
379 name: info.name,
380 backend,
381 adapter_type: format!("{:?}", info.device_type),
382 vendor: info.vendor,
383 device: info.device,
384 requested,
385 granted_limits: WgpuLimitEvidence::from_limits(&device.limits()),
386 granted_features: WgpuCapabilityEvidence::from_features(device.features()),
387 },
388 probe: ProbeEvidence {
389 transfer,
390 allocation_attempts,
391 },
392 },
393 device,
394 queue,
395 })
396}
397
398fn probe_transfer(
399 device: &wgpu::Device,
400 queue: &wgpu::Queue,
401 bytes: u64,
402) -> Result<TransferEvidence, WgpuDiscoveryError> {
403 let bytes = bytes.max(4).next_multiple_of(4);
404 let payload = (0..bytes).map(|idx| (idx % 251) as u8).collect::<Vec<_>>();
405 let source = device.create_buffer(&BufferDescriptor {
406 label: Some("sim-compute-wgpu-transfer-source"),
407 size: bytes,
408 usage: BufferUsages::COPY_SRC | BufferUsages::COPY_DST,
409 mapped_at_creation: false,
410 });
411 let readback = device.create_buffer(&BufferDescriptor {
412 label: Some("sim-compute-wgpu-transfer-readback"),
413 size: bytes,
414 usage: BufferUsages::COPY_DST | BufferUsages::MAP_READ,
415 mapped_at_creation: false,
416 });
417 queue.write_buffer(&source, 0, &payload);
418 let mut encoder = device.create_command_encoder(&wgpu::CommandEncoderDescriptor {
419 label: Some("sim-compute-wgpu-transfer-encoder"),
420 });
421 encoder.copy_buffer_to_buffer(&source, 0, &readback, 0, bytes);
422 queue.submit([encoder.finish()]);
423
424 let (sender, receiver) = mpsc::channel();
425 readback.slice(..).map_async(MapMode::Read, move |result| {
426 let _ = sender.send(result);
427 });
428 device
429 .poll(PollType::wait_indefinitely())
430 .map_err(|err| WgpuDiscoveryError::new(format!("wgpu poll failed: {err}")))?;
431 receiver
432 .recv()
433 .map_err(|err| WgpuDiscoveryError::new(format!("wgpu map callback failed: {err}")))?
434 .map_err(|err| WgpuDiscoveryError::new(format!("wgpu map failed: {err}")))?;
435
436 let mapped = readback
437 .slice(..)
438 .get_mapped_range()
439 .map_err(|err| WgpuDiscoveryError::new(format!("wgpu mapped range failed: {err}")))?
440 .to_vec();
441 readback.unmap();
442 Ok(TransferEvidence {
443 bytes,
444 transfer_ok: true,
445 mapping_ok: mapped == payload,
446 })
447}
448
449fn probe_allocations(device: &wgpu::Device, ceiling: u64) -> Vec<AllocationAttempt> {
450 [4096, 1024 * 1024, ceiling]
451 .into_iter()
452 .filter(|bytes| *bytes > 0)
453 .map(|bytes| {
454 let buffer = device.create_buffer(&BufferDescriptor {
455 label: Some("sim-compute-wgpu-allocation-probe"),
456 size: bytes,
457 usage: BufferUsages::COPY_DST,
458 mapped_at_creation: false,
459 });
460 drop(buffer);
461 AllocationAttempt {
462 bytes,
463 success: true,
464 }
465 })
466 .chain(std::iter::once(AllocationAttempt {
467 bytes: ceiling.saturating_add(1),
468 success: false,
469 }))
470 .collect()
471}