1use {
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#[derive(Clone)]
30pub struct ComputePipeline {
31 pub(crate) inner: Arc<ComputePipelineInner>,
32}
33
34impl ComputePipeline {
35 #[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 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 pub fn device(&self) -> &Device {
182 &self.inner.device
183 }
184
185 pub fn handle(&self) -> vk::Pipeline {
187 self.inner.handle
188 }
189
190 pub fn info(&self) -> ComputePipelineInfo {
192 self.inner.info
193 }
194
195 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 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#[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 #[builder(default = "8192")]
289 pub bindless_descriptor_count: u32,
290}
291
292impl ComputePipelineInfo {
293 pub fn builder() -> ComputePipelineInfoBuilder {
295 Default::default()
296 }
297
298 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 #[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}