Skip to main content

vk_graph/driver/
compute.rs

1//! Compute pipeline types.
2
3use {
4    super::{
5        DriverError,
6        device::Device,
7        shader::{DescriptorBindingMap, PipelineDescriptorInfo, Shader},
8    },
9    crate::{driver::DescriptorSetLayout, lazy_str},
10    ash::vk::{self, Handle as _},
11    derive_builder::Builder,
12    log::{trace, warn},
13    std::{
14        collections::HashSet,
15        ffi::CString,
16        fmt::{Debug, Formatter},
17        hash::{Hash, Hasher},
18        slice,
19        sync::Arc,
20        thread::panicking,
21    },
22};
23
24/// Smart pointer handle of a compute pipeline object.
25///
26/// Also contains information about the object.
27///
28/// See [`VkPipeline`](https://registry.khronos.org/vulkan/specs/latest/man/html/VkPipeline.html).
29#[derive(Clone)]
30pub struct ComputePipeline {
31    pub(crate) inner: Arc<ComputePipelineInner>,
32}
33
34impl ComputePipeline {
35    /// Creates a new compute pipeline on the given device.
36    ///
37    /// `shader` may be a pre-built [`Shader`] or any input that can be converted into one.
38    /// Invalid shader data is returned as [`DriverError::InvalidData`] through the `Result`
39    /// instead of panicking.
40    ///
41    /// See [`VkComputePipelineCreateInfo`](https://registry.khronos.org/vulkan/specs/latest/man/html/VkComputePipelineCreateInfo.html).
42    ///
43    /// # Examples
44    ///
45    /// Basic usage:
46    ///
47    /// ```no_run
48    /// # use std::sync::Arc;
49    /// # use ash::vk;
50    /// # use vk_graph::driver::DriverError;
51    /// # use vk_graph::driver::device::{Device, DeviceInfo};
52    /// # use vk_graph::driver::compute::{ComputePipeline, ComputePipelineInfo};
53    /// # use vk_graph::driver::shader::{Shader};
54    /// # fn main() -> Result<(), DriverError> {
55    /// # let device = Device::create(DeviceInfo::default())?;
56    /// # let my_shader_code = [0u8; 1];
57    /// // my_shader_code is raw SPIR-V code as bytes
58    /// let shader = Shader::new_compute(my_shader_code.as_slice());
59    /// let pipeline = ComputePipeline::create(&device, ComputePipelineInfo::default(), shader)?;
60    ///
61    /// assert_ne!(pipeline.handle(), vk::Pipeline::null());
62    /// # Ok(()) }
63    /// ```
64    #[profiling::function]
65    pub fn create<S>(
66        device: &Device,
67        info: impl Into<ComputePipelineInfo>,
68        shader: S,
69    ) -> Result<Self, DriverError>
70    where
71        S: TryInto<Shader>,
72        S::Error: Into<DriverError>,
73    {
74        trace!("create");
75
76        let info = info.into();
77        let shader = shader.try_into().map_err(Into::into)?;
78
79        // Use SPIR-V reflection to get the types and counts of all descriptors
80        let mut descriptor_bindings = shader.descriptor_bindings();
81        let mut bindless_descriptors = HashSet::new();
82        for (descriptor, (descriptor_info, _)) in descriptor_bindings.iter_mut() {
83            if descriptor_info.binding_count() == 0 {
84                bindless_descriptors.insert(*descriptor);
85                descriptor_info.set_binding_count(info.bindless_descriptor_count);
86            }
87        }
88
89        let descriptor_info =
90            PipelineDescriptorInfo::create(device, &descriptor_bindings, &bindless_descriptors)?;
91        let descriptor_set_layouts = descriptor_info
92            .layouts
93            .values()
94            .map(DescriptorSetLayout::handle)
95            .collect::<Box<_>>();
96
97        unsafe {
98            let shader_module = device
99                .create_shader_module(
100                    &vk::ShaderModuleCreateInfo::default().code(shader.spirv.words()),
101                    None,
102                )
103                .map_err(|err| {
104                    warn!("unable to create compute shader module: {err}");
105
106                    DriverError::Unsupported
107                })?;
108            let entry_name = CString::new(shader.entry_name.as_bytes()).map_err(|err| {
109                warn!("invalid compute shader entry name: {err}");
110
111                DriverError::InvalidData
112            })?;
113            let mut stage_create_info = vk::PipelineShaderStageCreateInfo::default()
114                .module(shader_module)
115                .stage(shader.stage)
116                .name(&entry_name);
117            let specialization_info = shader.specialization.as_ref().map(Into::into);
118
119            if let Some(specialization_info) = &specialization_info {
120                stage_create_info = stage_create_info.specialization_info(specialization_info);
121            }
122
123            let mut layout_info =
124                vk::PipelineLayoutCreateInfo::default().set_layouts(&descriptor_set_layouts);
125
126            let push_constants = shader.push_constant_range();
127            if let Some(push_constants) = &push_constants {
128                layout_info = layout_info.push_constant_ranges(slice::from_ref(push_constants));
129            }
130
131            let layout = device
132                .create_pipeline_layout(&layout_info, None)
133                .map_err(|err| {
134                    warn!("unable to create compute pipeline layout: {err}");
135
136                    device.destroy_shader_module(shader_module, None);
137
138                    DriverError::Unsupported
139                })?;
140            let create_info = vk::ComputePipelineCreateInfo::default()
141                .stage(stage_create_info)
142                .layout(layout);
143            let handle = device
144                .create_compute_pipelines(
145                    Device::pipeline_cache(device),
146                    slice::from_ref(&create_info),
147                    None,
148                )
149                .map_err(|(_, err)| {
150                    warn!("unable to create compute pipeline: {err}");
151
152                    device.destroy_shader_module(shader_module, None);
153
154                    DriverError::Unsupported
155                })?
156                .into_iter()
157                .find(|handle| !handle.is_null())
158                .ok_or_else(|| {
159                    warn!("missing pipeline handle");
160
161                    DriverError::Unsupported
162                })?;
163
164            device.destroy_shader_module(shader_module, None);
165
166            Ok(ComputePipeline {
167                inner: Arc::new(ComputePipelineInner {
168                    descriptor_bindings,
169                    descriptor_info,
170                    device: device.clone(),
171                    handle,
172                    info,
173                    layout,
174                    push_constants,
175                }),
176            })
177        }
178    }
179
180    /// The device which owns this compute pipeline.
181    pub fn device(&self) -> &Device {
182        &self.inner.device
183    }
184
185    /// The native Vulkan pipeline handle of this compute pipeline.
186    pub fn handle(&self) -> vk::Pipeline {
187        self.inner.handle
188    }
189
190    /// Gets the information used to create this object.
191    pub fn info(&self) -> ComputePipelineInfo {
192        self.inner.info
193    }
194
195    /// Sets the debugging name assigned to this pipeline.
196    pub fn set_debug_name(&self, name: impl AsRef<str>) {
197        Device::try_set_debug_utils_object_name(&self.inner.device, self.inner.handle, &name);
198        Device::try_set_private_data_object_name(
199            &self.inner.device,
200            vk::ObjectType::PIPELINE,
201            self.inner.handle,
202            &name,
203        );
204
205        Device::try_set_debug_utils_object_name(
206            &self.inner.device,
207            self.inner.layout,
208            lazy_str!("{} (layout)", name.as_ref()),
209        );
210
211        for (set_idx, layout) in &self.inner.descriptor_info.layouts {
212            layout.set_debug_name(lazy_str!("{} (DS{set_idx})", name.as_ref()));
213        }
214    }
215
216    /// Sets the debugging name assigned to this pipeline.
217    pub fn with_debug_name(self, name: impl AsRef<str>) -> Self {
218        self.set_debug_name(name);
219
220        self
221    }
222}
223
224impl Debug for ComputePipeline {
225    fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
226        let mut res = f.debug_struct(stringify!(ComputePipeline));
227
228        if let Some(debug_name) = &Device::private_data_object_name(
229            &self.inner.device,
230            vk::ObjectType::PIPELINE,
231            self.inner.handle,
232        ) {
233            res.field("debug_name", debug_name);
234        }
235
236        res.field("handle", &self.inner.handle)
237            .finish_non_exhaustive()
238    }
239}
240
241impl Eq for ComputePipeline {}
242
243impl Hash for ComputePipeline {
244    fn hash<H: Hasher>(&self, state: &mut H) {
245        Arc::as_ptr(&self.inner).hash(state);
246    }
247}
248
249impl PartialEq for ComputePipeline {
250    fn eq(&self, other: &Self) -> bool {
251        Arc::ptr_eq(&self.inner, &other.inner)
252    }
253}
254
255/// Information used to create a [`ComputePipeline`] instance.
256///
257/// See [`VkComputePipelineCreateInfo`](https://registry.khronos.org/vulkan/specs/latest/man/html/VkComputePipelineCreateInfo.html).
258#[derive(Builder, Clone, Copy, Debug, Eq, Hash, PartialEq)]
259#[builder(
260    build_fn(private, name = "fallible_build"),
261    derive(Clone, Copy, Debug),
262    pattern = "owned"
263)]
264pub struct ComputePipelineInfo {
265    /// The number of descriptors to allocate for a given binding when using bindless (unbounded)
266    /// syntax.
267    ///
268    /// The default is `8192`.
269    ///
270    /// # Examples
271    ///
272    /// Basic usage (GLSL):
273    ///
274    /// ```
275    /// # vk_shader_macros::glsl!(r#"
276    /// #version 460 core
277    /// #extension GL_EXT_nonuniform_qualifier : require
278    /// #pragma shader_stage(compute)
279    ///
280    /// layout(set = 0, binding = 0, rgba8) writeonly uniform image2D my_binding[];
281    ///
282    /// void main()
283    /// {
284    ///     // my_binding will have space for 8,192 images by default
285    /// }
286    /// # "#);
287    /// ```
288    #[builder(default = "8192")]
289    pub bindless_descriptor_count: u32,
290}
291
292impl ComputePipelineInfo {
293    /// Creates a default `ComputePipelineInfoBuilder`.
294    pub fn builder() -> ComputePipelineInfoBuilder {
295        Default::default()
296    }
297
298    /// Converts a `ComputePipelineInfo` into a `ComputePipelineInfoBuilder`.
299    pub fn into_builder(self) -> ComputePipelineInfoBuilder {
300        ComputePipelineInfoBuilder {
301            bindless_descriptor_count: Some(self.bindless_descriptor_count),
302        }
303    }
304}
305
306impl Default for ComputePipelineInfo {
307    fn default() -> Self {
308        Self {
309            bindless_descriptor_count: 8192,
310        }
311    }
312}
313
314impl From<ComputePipelineInfoBuilder> for ComputePipelineInfo {
315    fn from(info: ComputePipelineInfoBuilder) -> Self {
316        info.build()
317    }
318}
319
320impl ComputePipelineInfoBuilder {
321    /// Builds a new `ComputePipelineInfo`.
322    #[inline(always)]
323    pub fn build(self) -> ComputePipelineInfo {
324        self.fallible_build()
325            .expect("invalid compute pipeline info")
326    }
327}
328
329#[derive(Debug)]
330pub(crate) struct ComputePipelineInner {
331    pub descriptor_bindings: DescriptorBindingMap,
332    pub descriptor_info: PipelineDescriptorInfo,
333    pub device: Device,
334    pub handle: vk::Pipeline,
335    pub info: ComputePipelineInfo,
336    pub layout: vk::PipelineLayout,
337    pub push_constants: Option<vk::PushConstantRange>,
338}
339
340impl Drop for ComputePipelineInner {
341    #[profiling::function]
342    fn drop(&mut self) {
343        if panicking() {
344            return;
345        }
346
347        Device::try_clear_private_data_object_name(
348            &self.device,
349            vk::ObjectType::PIPELINE,
350            self.handle,
351        );
352
353        unsafe {
354            self.device.destroy_pipeline(self.handle, None);
355            self.device.destroy_pipeline_layout(self.layout, None);
356        }
357    }
358}
359
360#[cfg(test)]
361mod test {
362    use super::*;
363
364    type Info = ComputePipelineInfo;
365    type Builder = ComputePipelineInfoBuilder;
366
367    #[test]
368    pub fn compute_pipeline_info() {
369        let info = Info::default();
370        let builder = info.into_builder().build();
371
372        assert_eq!(info, builder);
373    }
374
375    #[test]
376    pub fn compute_pipeline_info_builder() {
377        let info = Info::default();
378        let builder = Builder::default().build();
379
380        assert_eq!(info, builder);
381    }
382}