1use std::{process::Command, sync::mpsc};
6
7use sim_lib_compute_auto::{ComputeDeviceIdentity, ComputeEvidenceKind};
8use sim_lib_compute_cuda::{
9 CudaLoadError, CudaProbePort, CudaRuntimeLoader, CudaRuntimeProbe, DynamicCudaLoader,
10};
11use sim_lib_compute_rocm::{
12 DynamicRocmLoader, RocmLoadError, RocmProbePort, RocmRuntimeLoader, RocmRuntimeProbe,
13};
14use sim_lib_compute_wgpu::{
15 AllocationAttempt, ProbeEvidence, ProbePolicy, RequestedWgpuProfile, TransferEvidence,
16 WgpuAdapterEvidence, WgpuAdapterProbe, WgpuAdapterRuntime, WgpuCapabilityEvidence,
17 WgpuDiscoveryError, WgpuLimitEvidence, WgpuProbePort,
18};
19use wgpu::{
20 BufferDescriptor, BufferUsages, DeviceDescriptor, ExperimentalFeatures, Features, Instance,
21 Limits, MapMode, MemoryHints, PollType,
22};
23
24#[derive(Clone, Debug)]
26pub struct UbuntuComputeProbe {
27 cuda: CudaRuntimeLoader,
28 rocm: RocmRuntimeLoader,
29}
30
31impl UbuntuComputeProbe {
32 #[must_use]
34 pub fn new() -> Self {
35 Self {
36 cuda: CudaRuntimeLoader::new(),
37 rocm: RocmRuntimeLoader::new().with_observed_gfx_targets(observed_gfx_targets()),
38 }
39 }
40}
41
42impl Default for UbuntuComputeProbe {
43 fn default() -> Self {
44 Self::new()
45 }
46}
47
48fn observed_gfx_targets() -> Vec<String> {
49 let Ok(output) = Command::new("rocm_agent_enumerator").output() else {
50 return Vec::new();
51 };
52 if !output.status.success() {
53 return Vec::new();
54 }
55 let Ok(stdout) = String::from_utf8(output.stdout) else {
56 return Vec::new();
57 };
58 parse_gfx_targets(&stdout)
59}
60
61fn parse_gfx_targets(stdout: &str) -> Vec<String> {
62 let mut targets = stdout
63 .lines()
64 .map(str::trim)
65 .filter(|line| line.starts_with("gfx") && *line != "gfx000")
66 .map(ToOwned::to_owned)
67 .collect::<Vec<_>>();
68 targets.sort();
69 targets.dedup();
70 targets
71}
72
73impl CudaProbePort for UbuntuComputeProbe {
74 fn probe_cuda(&self) -> Result<CudaRuntimeProbe, CudaLoadError> {
75 self.cuda.discover()
76 }
77}
78
79impl RocmProbePort for UbuntuComputeProbe {
80 fn probe_rocm(&self) -> Result<RocmRuntimeProbe, RocmLoadError> {
81 self.rocm.discover()
82 }
83}
84
85impl WgpuProbePort for UbuntuComputeProbe {
86 fn probe_wgpu(
87 &self,
88 policy: &ProbePolicy,
89 ) -> Result<Vec<WgpuAdapterRuntime>, WgpuDiscoveryError> {
90 let instance = Instance::default();
91 let adapters = pollster::block_on(instance.enumerate_adapters(policy.backends));
92 let mut runtimes = adapters
93 .into_iter()
94 .filter_map(|adapter| probe_adapter(&adapter, policy).ok())
95 .collect::<Vec<_>>();
96 runtimes.sort_by(|left, right| {
97 left.probe()
98 .adapter
99 .sort_key()
100 .cmp(&right.probe().adapter.sort_key())
101 });
102 Ok(runtimes)
103 }
104}
105
106fn probe_adapter(
107 adapter: &wgpu::Adapter,
108 policy: &ProbePolicy,
109) -> Result<WgpuAdapterRuntime, WgpuDiscoveryError> {
110 let info = adapter.get_info();
111 let backend = format!("{:?}", info.backend);
112 let identity = ComputeDeviceIdentity::new(info.name.clone(), "wgpu", backend.clone());
113 let required_features = adapter.features() & (Features::TIMESTAMP_QUERY | Features::SHADER_F16);
114 let required_limits = Limits::downlevel_defaults().using_resolution(adapter.limits());
115 let descriptor = DeviceDescriptor {
116 label: Some("sim-platform-ubuntu-compute-probe"),
117 required_features,
118 required_limits: required_limits.clone(),
119 experimental_features: ExperimentalFeatures::disabled(),
120 memory_hints: MemoryHints::Performance,
121 trace: wgpu::Trace::default(),
122 };
123 let (device, queue) = pollster::block_on(adapter.request_device(&descriptor))
124 .map_err(|error| WgpuDiscoveryError::new(format!("wgpu request_device failed: {error}")))?;
125 let bytes = policy.transfer_bytes.max(4).next_multiple_of(4);
126 let payload = (0..bytes)
127 .map(|index| (index % 251) as u8)
128 .collect::<Vec<_>>();
129 let source = device.create_buffer(&BufferDescriptor {
130 label: Some("sim-platform-wgpu-probe"),
131 size: bytes,
132 usage: BufferUsages::COPY_SRC | BufferUsages::COPY_DST,
133 mapped_at_creation: false,
134 });
135 let readback = device.create_buffer(&BufferDescriptor {
136 label: Some("sim-platform-wgpu-readback-probe"),
137 size: bytes,
138 usage: BufferUsages::COPY_DST | BufferUsages::MAP_READ,
139 mapped_at_creation: false,
140 });
141 queue.write_buffer(&source, 0, &payload);
142 let mut encoder = device.create_command_encoder(&wgpu::CommandEncoderDescriptor {
143 label: Some("sim-platform-wgpu-transfer-probe"),
144 });
145 encoder.copy_buffer_to_buffer(&source, 0, &readback, 0, bytes);
146 queue.submit([encoder.finish()]);
147 let (sender, receiver) = mpsc::channel();
148 readback.slice(..).map_async(MapMode::Read, move |result| {
149 let _ = sender.send(result);
150 });
151 device
152 .poll(PollType::wait_indefinitely())
153 .map_err(|error| WgpuDiscoveryError::new(format!("wgpu poll failed: {error}")))?;
154 receiver
155 .recv()
156 .map_err(|error| WgpuDiscoveryError::new(format!("wgpu map callback failed: {error}")))?
157 .map_err(|error| WgpuDiscoveryError::new(format!("wgpu map failed: {error}")))?;
158 let mapping_ok = readback
159 .slice(..)
160 .get_mapped_range()
161 .map_err(|error| WgpuDiscoveryError::new(format!("wgpu mapped range failed: {error}")))?
162 .as_ref()
163 == payload.as_slice();
164 readback.unmap();
165 let ceiling = device
166 .limits()
167 .max_buffer_size
168 .min(policy.max_allocation_probe_bytes)
169 .max(4);
170 let allocation = device.create_buffer(&BufferDescriptor {
171 label: Some("sim-platform-wgpu-allocation-probe"),
172 size: ceiling,
173 usage: BufferUsages::COPY_DST,
174 mapped_at_creation: false,
175 });
176 drop(allocation);
177 let probe = WgpuAdapterProbe {
178 evidence_kind: ComputeEvidenceKind::PhysicalDevice,
179 claimed_identity: Some(identity.clone()),
180 observed_identity: Some(identity),
181 adapter: WgpuAdapterEvidence {
182 ordinal: 0,
183 name: info.name,
184 backend,
185 adapter_type: format!("{:?}", info.device_type),
186 vendor: info.vendor,
187 device: info.device,
188 requested: RequestedWgpuProfile::from_parts(required_limits, required_features),
189 granted_limits: WgpuLimitEvidence::from_limits(&device.limits()),
190 granted_features: WgpuCapabilityEvidence::from_features(device.features()),
191 },
192 probe: ProbeEvidence {
193 transfer: TransferEvidence {
194 bytes,
195 transfer_ok: true,
196 mapping_ok,
197 },
198 allocation_attempts: vec![AllocationAttempt {
199 bytes: ceiling,
200 success: true,
201 }],
202 },
203 };
204 Ok(WgpuAdapterRuntime::new(probe, device, queue))
205}
206
207#[cfg(test)]
208mod tests {
209 use super::*;
210 use sim_lib_compute_cuda::ComputeCudaLib;
211 use sim_lib_compute_rocm::ComputeRocmLib;
212 use sim_lib_compute_wgpu::ComputeWgpuLib;
213
214 struct EmptyWgpuProbe;
215
216 impl WgpuProbePort for EmptyWgpuProbe {
217 fn probe_wgpu(
218 &self,
219 _policy: &ProbePolicy,
220 ) -> Result<Vec<WgpuAdapterRuntime>, WgpuDiscoveryError> {
221 Ok(Vec::new())
222 }
223 }
224
225 #[test]
226 fn amd_target_observation_is_filtered_sorted_and_deduplicated() {
227 assert_eq!(
228 parse_gfx_targets("gfx1151\ngfx000\nnoise\ngfx1103\ngfx1151\n"),
229 vec!["gfx1103", "gfx1151"]
230 );
231 }
232
233 #[test]
234 fn absent_native_runtimes_export_no_provider() {
235 let probe = UbuntuComputeProbe {
236 cuda: CudaRuntimeLoader::with_search_dirs_only(Vec::new()),
237 rocm: RocmRuntimeLoader::with_search_dirs_only(Vec::new()),
238 };
239 assert!(ComputeCudaLib::from_probe_port(&probe).is_err());
240 assert!(ComputeRocmLib::from_probe_port(&probe).is_err());
241 let wgpu =
242 ComputeWgpuLib::from_probe_port(&EmptyWgpuProbe, &ProbePolicy::default()).unwrap();
243 assert!(sim_kernel::Lib::manifest(&wgpu).exports.is_empty());
244 }
245}