axiolid_backend_gpu/
adapter.rs1use axiolid_contracts::{
13 Backend, BackendDescriptor, BackendId, DevicePreference, ExecutionOptions, ExecutionTarget,
14 GeomError, GeomResult, Operation, Precision, Residency,
15};
16use axiolid_mesh::TriMesh;
17use axiolid_mesh_compile_contract::MeshCompiler;
18use axiolid_model::{GeometryGraph, NodeId};
19
20use crate::{GpuDeviceDescriptor, GpuGraphExecutor};
21
22#[derive(Debug)]
24pub struct GpuCompiler<E> {
25 executor: E,
26}
27
28impl<E> GpuCompiler<E> {
29 pub const fn new(executor: E) -> Self {
31 Self { executor }
32 }
33
34 pub fn device(&self) -> &GpuDeviceDescriptor
36 where
37 E: GpuGraphExecutor,
38 {
39 self.executor.device()
40 }
41
42 pub const fn executor(&self) -> &E {
44 &self.executor
45 }
46
47 fn validate_options(&self, options: &ExecutionOptions) -> GeomResult<()>
48 where
49 E: GpuGraphExecutor,
50 {
51 let device = self.executor.device();
52 let compatible_device = match options.device() {
53 DevicePreference::Auto | DevicePreference::Gpu => true,
54 DevicePreference::Backend(required) => required == device.id,
55 DevicePreference::Cpu => false,
56 };
57 if !compatible_device || (options.precision() == Precision::F64 && !device.features.float64)
58 {
59 return Err(GeomError::Unsupported {
60 backend: device.id,
61 operation: Operation::GraphCompilation,
62 });
63 }
64 let output = options.residency().output();
69 let deliverable = match output {
70 Residency::Host => true,
71 Residency::Device(owner) | Residency::Unified(owner) => owner == device.id,
72 _ => false,
75 };
76 if !deliverable {
77 return Err(GeomError::Unsupported {
78 backend: device.id,
79 operation: Operation::GraphCompilation,
80 });
81 }
82 self.executor.validate_options(options)
83 }
84
85 fn validate_roots(graph: &GeometryGraph, roots: &[NodeId]) -> GeomResult<()> {
86 if let Some(root) = roots.iter().find(|root| graph.get(**root).is_none()) {
87 return Err(GeomError::InvalidInput(format!(
88 "graph compilation root {root} does not belong to the supplied graph"
89 )));
90 }
91 Ok(())
92 }
93
94 fn validate_result_count(
95 backend: BackendId,
96 root_count: usize,
97 result_count: usize,
98 ) -> GeomResult<()> {
99 if result_count != root_count {
100 return Err(GeomError::BackendContractViolation {
101 backend,
102 detail: format!("returned {result_count} results for {root_count} roots"),
103 });
104 }
105 Ok(())
106 }
107}
108
109impl<E: GpuGraphExecutor> Backend for GpuCompiler<E> {
110 fn descriptor(&self) -> BackendDescriptor {
111 BackendDescriptor::new(self.executor.device().id, ExecutionTarget::Gpu)
112 }
113}
114
115impl<E: GpuGraphExecutor> MeshCompiler for GpuCompiler<E> {
116 fn compile_mesh(
117 &self,
118 graph: &GeometryGraph,
119 root: NodeId,
120 options: &ExecutionOptions,
121 ) -> GeomResult<TriMesh> {
122 self.validate_options(options)?;
123 Self::validate_roots(graph, &[root])?;
124 let results = self.executor.compile_mesh_batch(graph, &[root], options)?;
125 Self::validate_result_count(self.executor.device().id, 1, results.len())?;
126 results
127 .into_iter()
128 .next()
129 .ok_or_else(|| GeomError::BackendContractViolation {
130 backend: self.executor.device().id,
131 detail: "returned no result for one root".to_owned(),
132 })
133 }
134
135 fn compile_mesh_batch_into(
139 &self,
140 graph: &GeometryGraph,
141 roots: &[NodeId],
142 options: &ExecutionOptions,
143 destination: &mut Vec<TriMesh>,
144 ) -> GeomResult<()> {
145 self.validate_options(options)?;
146 Self::validate_roots(graph, roots)?;
147 let results = self.executor.compile_mesh_batch(graph, roots, options)?;
148 Self::validate_result_count(self.executor.device().id, roots.len(), results.len())?;
149 destination.reserve(results.len());
150 destination.extend(results);
151 Ok(())
152 }
153}