1use std::cell::{Cell, RefCell};
13use std::collections::HashMap;
14use std::ffi::CStr;
15use std::ptr::NonNull;
16use std::rc::Rc;
17use std::sync::{Arc, Mutex, OnceLock};
18
19use ash::vk;
20use virtio_accel_core::{
21 Accelerator, AcceleratorClass, AccessMode, AllocatedBuffer, ArtifactRef, BackendError,
22 BindingRef, BufferDesc, BufferInfo, BufferProperties, BufferUsage, ByteSink, ByteSource,
23 Capabilities, ContextDesc, DeviceIdentity, DeviceInfo, DeviceLimits, EventState, MemoryDomain,
24 QueueDesc, ReleaseFailure, SubmitFailure, Timeout,
25};
26use virtio_accel_tosa::{CapabilityDescriptor, TosaCapabilityProvider};
27
28use crate::lower::{KernelSpec, LoweringError, ProgramPlan, SlotRole, Work, lower_tosa};
29use crate::nvfp4::lower_nvfp4;
30use crate::shader::{self, KernelKey};
31use crate::{InitError, REQUIRED_RESIDENT_BYTES};
32
33const MAX_TOSA_ARTIFACT_BYTES: u64 = 256 * 1024 * 1024;
35
36const VULKAN_EXTERNAL_DOMAIN: u32 = 0x5655_4c4b;
38
39const STAGING_BYTES: u64 = 4 * 1024 * 1024;
41
42const TRANSFER_TIMEOUT_NS: u64 = 30_000_000_000;
44
45const ASSUMED_MAX_MEMORY_ALLOCATIONS: u32 = 4096;
49const MAX_CONTEXTS: u32 = 16;
50const MAX_BUFFERS_PER_CONTEXT: u32 = 190;
51const MAX_PROGRAMS_PER_CONTEXT: u32 = 64;
53const MAX_QUEUES_PER_CONTEXT: u32 = 16;
54const RING_DEPTH: u32 = 64;
57const MAX_BINDINGS_PER_SUBMISSION: u32 = 16;
58
59const TRANSIENT_STAGING_ALLOCATIONS: u32 = 1;
62
63const _: () = assert!(
64 MAX_CONTEXTS * (MAX_BUFFERS_PER_CONTEXT + MAX_PROGRAMS_PER_CONTEXT)
65 + TRANSIENT_STAGING_ALLOCATIONS
66 <= ASSUMED_MAX_MEMORY_ALLOCATIONS,
67 "advertised buffers, program arenas, and the staging allocation must fit the assumed allocation count"
68);
69
70const PREFERRED_WORKGROUP: u32 = 256;
73const FALLBACK_WORKGROUP: u32 = 128;
74const PREFERRED_MATMUL_TILE: u32 = 16;
78const FALLBACK_MATMUL_TILE: u32 = 8;
79const WORD_BYTES: u64 = 4;
82
83#[derive(Clone, Copy, Debug, PartialEq, Eq)]
85struct Tuning {
86 workgroup: u32,
88 matmul_tile: u32,
90 buffers: u32,
92 cooperative_nvfp4: bool,
95 subgroup_nvfp4: bool,
98}
99
100impl Tuning {
101 fn from_limits(limits: &vk::PhysicalDeviceLimits) -> Option<Self> {
103 let invocations = limits.max_compute_work_group_invocations;
104 let size = limits.max_compute_work_group_size;
105 let workgroup = if invocations >= PREFERRED_WORKGROUP && size[0] >= PREFERRED_WORKGROUP {
106 PREFERRED_WORKGROUP
107 } else if invocations >= FALLBACK_WORKGROUP && size[0] >= FALLBACK_WORKGROUP {
108 FALLBACK_WORKGROUP
109 } else {
110 return None;
111 };
112 let tile_fits = |tile: u32| {
115 invocations >= tile * tile
116 && size[0] >= tile.max(shader::STREAM_WORKGROUP)
117 && size[1] >= tile
118 && limits.max_compute_shared_memory_size >= shader::matmul_shared_bytes(tile)
119 };
120 let matmul_tile = if tile_fits(PREFERRED_MATMUL_TILE) {
121 PREFERRED_MATMUL_TILE
122 } else if tile_fits(FALLBACK_MATMUL_TILE) {
123 FALLBACK_MATMUL_TILE
124 } else {
125 return None;
126 };
127 let descriptors = limits
130 .max_per_stage_descriptor_storage_buffers
131 .min(limits.max_descriptor_set_storage_buffers)
132 .min(MAX_BINDINGS_PER_SUBMISSION + 1);
133 if descriptors < 3 {
134 return None;
135 }
136 Some(Self {
137 workgroup,
138 matmul_tile,
139 buffers: descriptors,
140 cooperative_nvfp4: false,
141 subgroup_nvfp4: false,
142 })
143 }
144
145 const fn max_bindings(self) -> u32 {
147 self.buffers - 1
148 }
149
150 fn key(self, kernel: KernelSpec) -> KernelKey {
151 match kernel {
152 KernelSpec::Nvfp4Matmul { cooperative } => KernelKey::Nvfp4Matmul {
153 buffers: self.buffers,
154 cooperative: cooperative && self.cooperative_nvfp4,
155 subgroup: !(cooperative && self.cooperative_nvfp4) && self.subgroup_nvfp4,
156 },
157 KernelSpec::Elementwise {
158 op,
159 float,
160 broadcast,
161 } => KernelKey::Elementwise {
162 op,
163 float,
164 broadcast,
165 workgroup: self.workgroup,
166 buffers: self.buffers,
167 },
168 KernelSpec::Reduce { op, float } => KernelKey::Reduce {
169 op,
170 float,
171 workgroup: self.workgroup,
172 buffers: self.buffers,
173 },
174 KernelSpec::Matmul { input, output } => KernelKey::Matmul {
175 input,
176 output,
177 tile: self.matmul_tile,
178 buffers: self.buffers,
179 },
180 KernelSpec::MatmulStream { rhs, output } => KernelKey::MatmulStream {
181 rhs,
182 output,
183 buffers: self.buffers,
184 },
185 KernelSpec::MaxPool { nan_mode, float } => KernelKey::MaxPool {
186 nan_mode,
187 float,
188 workgroup: self.workgroup,
189 buffers: self.buffers,
190 },
191 KernelSpec::Cast { input, output } => KernelKey::Cast {
192 input,
193 output,
194 workgroup: self.workgroup,
195 buffers: self.buffers,
196 },
197 KernelSpec::Move {
198 storage,
199 contiguous,
200 } => KernelKey::Move {
201 storage,
202 contiguous,
203 workgroup: self.workgroup,
204 buffers: self.buffers,
205 },
206 }
207 }
208
209 fn workgroups(
211 self,
212 work: Work,
213 kernel: KernelSpec,
214 limits: &vk::PhysicalDeviceLimits,
215 ) -> Option<[u32; 3]> {
216 let max = limits.max_compute_work_group_count;
217 match work {
218 Work::Linear(count) => Some([
219 shader::linear_workgroups(count, self.workgroup, max[0].max(1)),
220 1,
221 1,
222 ]),
223 Work::Matmul { m, n, batch } => {
224 let groups = shader::matmul_workgroups(m, n, batch, self.matmul_tile);
225 (groups[0] <= max[0] && groups[1] <= max[1] && groups[2] <= max[2])
226 .then_some(groups)
227 }
228 Work::MatmulStream { n, batch } => {
229 let groups = shader::stream_matmul_workgroups(n, batch);
230 (groups[0] <= max[0] && groups[2] <= max[2]).then_some(groups)
231 }
232 Work::Nvfp4Matmul { m, n } => {
233 let cooperative = matches!(kernel, KernelSpec::Nvfp4Matmul { cooperative: true })
234 && self.cooperative_nvfp4;
235 let subgroup = !cooperative && self.subgroup_nvfp4;
236 let columns = if cooperative {
237 n.div_ceil(16)
238 } else if subgroup {
239 n.div_ceil(4)
240 } else {
241 n
242 };
243 let rows = if cooperative { m.div_ceil(8) } else { m };
244 let groups = [columns, rows, 1];
245 (groups[0] <= max[0] && groups[1] <= max[1]).then_some(groups)
246 }
247 }
248 }
249}
250
251const EXCLUSIVE_ACCESS: u64 = 1 << 63;
252
253fn entry() -> Result<ash::Entry, InitError> {
255 static ENTRY: OnceLock<Result<ash::Entry, InitError>> = OnceLock::new();
256 ENTRY
257 .get_or_init(|| {
258 unsafe { ash::Entry::load() }.map_err(|_| InitError::RuntimeUnavailable)
261 })
262 .clone()
263}
264
265fn backend_error(result: vk::Result) -> BackendError {
266 match result {
267 vk::Result::ERROR_OUT_OF_HOST_MEMORY | vk::Result::ERROR_OUT_OF_DEVICE_MEMORY => {
268 BackendError::OutOfMemory
269 }
270 vk::Result::ERROR_TOO_MANY_OBJECTS => BackendError::ResourceLimit,
271 vk::Result::ERROR_DEVICE_LOST => BackendError::DeviceLost,
272 vk::Result::ERROR_MEMORY_MAP_FAILED => BackendError::Incompatible,
273 other => BackendError::External {
274 domain: VULKAN_EXTERNAL_DOMAIN,
275 code: i64::from(other.as_raw()),
276 },
277 }
278}
279
280struct Instance {
282 _entry: ash::Entry,
284 instance: ash::Instance,
285}
286
287impl Instance {
288 fn create() -> Result<Self, InitError> {
289 let entry = entry()?;
290 let loader_version = unsafe { entry.try_enumerate_instance_version() }
292 .ok()
293 .flatten()
294 .unwrap_or(vk::API_VERSION_1_0);
295 if loader_version < vk::API_VERSION_1_3 {
296 return Err(InitError::RuntimeUnavailable);
297 }
298 let application = vk::ApplicationInfo::default()
299 .application_name(c"virtio-accel-vulkan")
300 .api_version(vk::API_VERSION_1_3);
301 let extension_names = [c"VK_KHR_portability_enumeration".as_ptr()];
304 let info = vk::InstanceCreateInfo::default()
305 .application_info(&application)
306 .flags(vk::InstanceCreateFlags::ENUMERATE_PORTABILITY_KHR)
307 .enabled_extension_names(&extension_names);
308 let instance = match unsafe { entry.create_instance(&info, None) } {
311 Ok(instance) => instance,
312 Err(vk::Result::ERROR_INCOMPATIBLE_DRIVER) => {
313 return Err(InitError::RuntimeUnavailable);
314 }
315 Err(_) => return Err(InitError::InstanceCreationFailed),
316 };
317 Ok(Self {
318 _entry: entry,
319 instance,
320 })
321 }
322}
323
324impl Drop for Instance {
325 fn drop(&mut self) {
326 unsafe { self.instance.destroy_instance(None) };
329 }
330}
331
332#[derive(Clone)]
334struct PhysicalDeviceRecord {
335 handle: vk::PhysicalDevice,
336 name: String,
337 device_type: vk::PhysicalDeviceType,
338 vendor_id: u32,
339 device_id: u32,
340 uuid: [u8; 16],
341 queue_family: u32,
342 limits: vk::PhysicalDeviceLimits,
343 memory: vk::PhysicalDeviceMemoryProperties,
344 buffer_device_address: bool,
345 host_import_alignment: Option<u64>,
349 timeline_semaphore: bool,
351 shader_float16: bool,
352 vulkan_memory_model: bool,
353 cooperative_nvfp4: bool,
354 tuning: Tuning,
355}
356
357impl PhysicalDeviceRecord {
358 fn probe(owner: &Instance, handle: vk::PhysicalDevice) -> Option<Self> {
361 let instance = &owner.instance;
362 let mut vulkan11 = vk::PhysicalDeviceVulkan11Properties::default();
363 let mut subgroup = vk::PhysicalDeviceSubgroupProperties::default();
364 let mut properties = vk::PhysicalDeviceProperties2::default()
365 .push_next(&mut vulkan11)
366 .push_next(&mut subgroup);
367 unsafe { instance.get_physical_device_properties2(handle, &mut properties) };
369 let properties = properties.properties;
370 if properties.api_version < vk::API_VERSION_1_3 {
371 return None;
372 }
373 let mut tuning = Tuning::from_limits(&properties.limits)?;
374
375 let mut vulkan12 = vk::PhysicalDeviceVulkan12Features::default();
376 let mut vulkan13 = vk::PhysicalDeviceVulkan13Features::default();
377 let mut cooperative = vk::PhysicalDeviceCooperativeMatrixFeaturesKHR::default();
378 let mut features = vk::PhysicalDeviceFeatures2::default()
379 .push_next(&mut vulkan12)
380 .push_next(&mut vulkan13)
381 .push_next(&mut cooperative);
382 unsafe { instance.get_physical_device_features2(handle, &mut features) };
384 if vulkan13.synchronization2 == vk::FALSE {
385 return None;
386 }
387
388 let families = unsafe { instance.get_physical_device_queue_family_properties(handle) };
390 let compute_only = families.iter().position(|family| {
393 family.queue_flags.contains(vk::QueueFlags::COMPUTE)
394 && !family.queue_flags.contains(vk::QueueFlags::GRAPHICS)
395 });
396 let any_compute = families
397 .iter()
398 .position(|family| family.queue_flags.contains(vk::QueueFlags::COMPUTE));
399 let queue_family = u32::try_from(compute_only.or(any_compute)?).ok()?;
400
401 let memory = unsafe { instance.get_physical_device_memory_properties(handle) };
403 let host_import_alignment = host_import_alignment(instance, handle);
404 let cooperative_nvfp4 = vulkan12.shader_float16 == vk::TRUE
405 && vulkan12.vulkan_memory_model == vk::TRUE
406 && cooperative.cooperative_matrix == vk::TRUE
407 && subgroup.subgroup_size == 32
408 && subgroup
409 .supported_stages
410 .contains(vk::ShaderStageFlags::COMPUTE)
411 && supports_extension(instance, handle, ash::khr::cooperative_matrix::NAME)
412 && supports_nvfp4_cooperative_matrix(owner, handle);
413 tuning.cooperative_nvfp4 = cooperative_nvfp4;
414 tuning.subgroup_nvfp4 = subgroup.subgroup_size == 32
415 && subgroup
416 .supported_stages
417 .contains(vk::ShaderStageFlags::COMPUTE)
418 && subgroup
419 .supported_operations
420 .contains(vk::SubgroupFeatureFlags::ARITHMETIC);
421 let name = CStr::from_bytes_until_nul(bytemuck_i8_to_u8(&properties.device_name))
422 .map(|name| name.to_string_lossy().into_owned())
423 .unwrap_or_else(|_| String::from("vulkan-device"));
424 Some(Self {
425 handle,
426 name,
427 device_type: properties.device_type,
428 vendor_id: properties.vendor_id,
429 device_id: properties.device_id,
430 uuid: vulkan11.device_uuid,
431 queue_family,
432 limits: properties.limits,
433 memory,
434 buffer_device_address: vulkan12.buffer_device_address == vk::TRUE,
435 host_import_alignment,
436 timeline_semaphore: vulkan12.timeline_semaphore == vk::TRUE,
437 shader_float16: vulkan12.shader_float16 == vk::TRUE,
438 vulkan_memory_model: vulkan12.vulkan_memory_model == vk::TRUE,
439 cooperative_nvfp4,
440 tuning,
441 })
442 }
443
444 fn rank(&self) -> u8 {
446 match self.device_type {
447 vk::PhysicalDeviceType::DISCRETE_GPU => 0,
448 vk::PhysicalDeviceType::INTEGRATED_GPU => 1,
449 vk::PhysicalDeviceType::VIRTUAL_GPU => 2,
450 vk::PhysicalDeviceType::CPU => 3,
451 _ => 4,
452 }
453 }
454
455 fn class(&self) -> AcceleratorClass {
456 if self.device_type == vk::PhysicalDeviceType::CPU {
457 AcceleratorClass::OTHER
458 } else {
459 AcceleratorClass::GPU
460 }
461 }
462}
463
464fn host_import_alignment(instance: &ash::Instance, handle: vk::PhysicalDevice) -> Option<u64> {
466 let extensions = unsafe { instance.enumerate_device_extension_properties(handle) }.ok()?;
468 let name = ash::ext::external_memory_host::NAME;
469 extensions
470 .iter()
471 .any(|extension| extension.extension_name_as_c_str() == Ok(name))
472 .then_some(())?;
473 let mut host = vk::PhysicalDeviceExternalMemoryHostPropertiesEXT::default();
474 let mut properties = vk::PhysicalDeviceProperties2::default().push_next(&mut host);
475 unsafe { instance.get_physical_device_properties2(handle, &mut properties) };
477 let alignment = host.min_imported_host_pointer_alignment;
478 alignment.is_power_of_two().then_some(alignment)
479}
480
481fn supports_extension(instance: &ash::Instance, handle: vk::PhysicalDevice, name: &CStr) -> bool {
482 unsafe { instance.enumerate_device_extension_properties(handle) }.is_ok_and(|extensions| {
484 extensions
485 .iter()
486 .any(|extension| extension.extension_name_as_c_str() == Ok(name))
487 })
488}
489
490fn supports_nvfp4_cooperative_matrix(owner: &Instance, handle: vk::PhysicalDevice) -> bool {
491 let extension = ash::khr::cooperative_matrix::Instance::new(&owner._entry, &owner.instance);
492 let Ok(properties) =
495 (unsafe { extension.get_physical_device_cooperative_matrix_properties(handle) })
496 else {
497 return false;
498 };
499 properties.iter().any(|property| {
500 property.scope == vk::ScopeKHR::SUBGROUP
501 && property.m_size == 8
502 && property.n_size == 16
503 && property.k_size == 16
504 && property.a_type == vk::ComponentTypeKHR::FLOAT16
505 && property.b_type == vk::ComponentTypeKHR::FLOAT16
506 && property.c_type == vk::ComponentTypeKHR::FLOAT32
507 && property.result_type == vk::ComponentTypeKHR::FLOAT32
508 && property.saturating_accumulation == vk::FALSE
509 })
510}
511
512fn bytemuck_i8_to_u8(name: &[std::ffi::c_char; 256]) -> &[u8; 256] {
514 unsafe { &*(name as *const [std::ffi::c_char; 256]).cast::<[u8; 256]>() }
516}
517
518fn enumerate(instance: &Instance) -> Result<Vec<PhysicalDeviceRecord>, InitError> {
519 let handles = unsafe { instance.instance.enumerate_physical_devices() }
521 .map_err(|_| InitError::DeviceEnumerationFailed)?;
522 Ok(handles
523 .into_iter()
524 .filter_map(|handle| PhysicalDeviceRecord::probe(instance, handle))
525 .collect())
526}
527
528#[derive(Clone, Copy, Debug)]
530struct MemoryPlan {
531 host: u32,
533 device: Option<u32>,
535 shared: Option<u32>,
537}
538
539#[derive(Clone, Copy, Debug, PartialEq, Eq)]
542pub struct VulkanOptions {
543 pub map_unified_device_memory: bool,
550}
551
552impl Default for VulkanOptions {
553 fn default() -> Self {
554 Self {
555 map_unified_device_memory: true,
556 }
557 }
558}
559
560impl MemoryPlan {
561 fn select(
565 memory: &vk::PhysicalDeviceMemoryProperties,
566 buffer_type_mask: u32,
567 unified: bool,
568 ) -> Option<Self> {
569 let types = &memory.memory_types[..memory.memory_type_count as usize];
570 let usable = |index: usize, flags: vk::MemoryPropertyFlags| {
571 buffer_type_mask & (1 << index) != 0
572 && !flags.intersects(
573 vk::MemoryPropertyFlags::PROTECTED
574 | vk::MemoryPropertyFlags::LAZILY_ALLOCATED
575 | vk::MemoryPropertyFlags::DEVICE_COHERENT_AMD
576 | vk::MemoryPropertyFlags::RDMA_CAPABLE_NV,
577 )
578 };
579 let pick = |required: vk::MemoryPropertyFlags,
580 score: fn(vk::MemoryPropertyFlags) -> u8|
581 -> Option<u32> {
582 types
583 .iter()
584 .enumerate()
585 .filter(|(index, ty)| {
586 usable(*index, ty.property_flags) && ty.property_flags.contains(required)
587 })
588 .max_by_key(|(_, ty)| score(ty.property_flags))
589 .and_then(|(index, _)| u32::try_from(index).ok())
590 };
591 let host_coherent =
592 vk::MemoryPropertyFlags::HOST_VISIBLE | vk::MemoryPropertyFlags::HOST_COHERENT;
593 Some(Self {
594 host: pick(host_coherent, |flags| {
595 u8::from(!flags.contains(vk::MemoryPropertyFlags::DEVICE_LOCAL)) * 2
596 + u8::from(flags.contains(vk::MemoryPropertyFlags::HOST_CACHED))
597 })?,
598 device: pick(
599 vk::MemoryPropertyFlags::DEVICE_LOCAL,
600 if unified {
601 |flags| {
602 u8::from(flags.contains(
603 vk::MemoryPropertyFlags::HOST_VISIBLE
604 | vk::MemoryPropertyFlags::HOST_COHERENT,
605 )) * 2
606 + u8::from(flags.contains(vk::MemoryPropertyFlags::HOST_CACHED))
607 }
608 } else {
609 |flags| u8::from(!flags.contains(vk::MemoryPropertyFlags::HOST_VISIBLE))
610 },
611 ),
612 shared: pick(
613 vk::MemoryPropertyFlags::DEVICE_LOCAL | host_coherent,
614 |flags| u8::from(flags.contains(vk::MemoryPropertyFlags::HOST_CACHED)),
615 ),
616 })
617 }
618
619 fn is_mapped(self, memory: &vk::PhysicalDeviceMemoryProperties, domain: MemoryDomain) -> bool {
622 self.for_domain(domain).is_some_and(|index| {
623 memory.memory_types[index as usize].property_flags.contains(
624 vk::MemoryPropertyFlags::HOST_VISIBLE | vk::MemoryPropertyFlags::HOST_COHERENT,
625 )
626 })
627 }
628
629 fn for_domain(self, domain: MemoryDomain) -> Option<u32> {
630 match domain {
631 MemoryDomain::Host => Some(self.host),
632 MemoryDomain::Device => self.device,
633 MemoryDomain::Shared => self.shared,
634 }
635 }
636
637 fn capabilities(self) -> Capabilities {
638 let mut capabilities = Capabilities::HOST_VISIBLE_MEMORY;
639 if self.device.is_some() {
640 capabilities |= Capabilities::DEVICE_LOCAL_MEMORY;
641 }
642 if self.shared.is_some() {
643 capabilities |= Capabilities::SHARED_MEMORY;
644 }
645 capabilities
646 }
647}
648
649#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
651pub struct LiveResources {
652 pub contexts: u64,
653 pub buffers: u64,
654 pub programs: u64,
655 pub queues: u64,
656 pub events: u64,
657}
658
659#[derive(Default)]
660struct Counters {
661 direct_binding_admissions: Cell<u64>,
662 explicit_transfer_bytes: Cell<u64>,
663 contexts: Cell<u64>,
664 buffers: Cell<u64>,
665 programs: Cell<u64>,
666 queues: Cell<u64>,
667 events: Cell<u64>,
668}
669
670fn increment(cell: &Cell<u64>, by: u64) {
671 cell.set(cell.get().saturating_add(by));
672}
673
674fn decrement(cell: &Cell<u64>) {
675 cell.set(cell.get().saturating_sub(1));
676}
677
678struct Shared {
683 device: ash::Device,
684 physical: PhysicalDeviceRecord,
685 queue: vk::Queue,
686 set_layout: vk::DescriptorSetLayout,
687 pipeline_layout: vk::PipelineLayout,
688 pipeline_cache: vk::PipelineCache,
691 modules: RefCell<HashMap<KernelKey, Rc<[u32]>>>,
693 memory_plan: MemoryPlan,
694 host_memory: Option<ash::ext::external_memory_host::Device>,
696 info: DeviceInfo,
697 poisoned: Cell<bool>,
700 counters: Counters,
701 _instance: Instance,
703}
704
705impl Shared {
706 fn open(
707 instance: Instance,
708 physical: PhysicalDeviceRecord,
709 options: VulkanOptions,
710 ) -> Result<Rc<Self>, InitError> {
711 let priorities = [1.0_f32];
712 let queue_info = vk::DeviceQueueCreateInfo::default()
713 .queue_family_index(physical.queue_family)
714 .queue_priorities(&priorities);
715 let queue_infos = [queue_info];
716 let mut vulkan12 = vk::PhysicalDeviceVulkan12Features::default()
719 .buffer_device_address(physical.buffer_device_address)
720 .timeline_semaphore(physical.timeline_semaphore)
721 .shader_float16(physical.shader_float16)
722 .vulkan_memory_model(physical.vulkan_memory_model);
723 let mut vulkan13 = vk::PhysicalDeviceVulkan13Features::default().synchronization2(true);
724 let mut cooperative = vk::PhysicalDeviceCooperativeMatrixFeaturesKHR::default()
725 .cooperative_matrix(physical.cooperative_nvfp4);
726 let mut extensions = Vec::with_capacity(2);
729 if physical.host_import_alignment.is_some() {
730 extensions.push(ash::ext::external_memory_host::NAME.as_ptr());
731 }
732 if physical.cooperative_nvfp4 {
733 extensions.push(ash::khr::cooperative_matrix::NAME.as_ptr());
734 }
735 let info = vk::DeviceCreateInfo::default()
736 .queue_create_infos(&queue_infos)
737 .enabled_extension_names(&extensions)
738 .push_next(&mut vulkan12)
739 .push_next(&mut vulkan13)
740 .push_next(&mut cooperative);
741 let device = unsafe {
745 instance
746 .instance
747 .create_device(physical.handle, &info, None)
748 }
749 .map_err(|_| InitError::DeviceCreationFailed)?;
750 let queue = unsafe { device.get_device_queue(physical.queue_family, 0) };
752
753 let buffer_type_mask = match probe_buffer_type_mask(&device, &physical) {
757 Ok(mask) => mask,
758 Err(_) => {
759 unsafe { device.destroy_device(None) };
761 return Err(InitError::DeviceCreationFailed);
762 }
763 };
764 let unified = options.map_unified_device_memory && physical.memory.memory_heap_count == 1;
765 let Some(memory_plan) = MemoryPlan::select(&physical.memory, buffer_type_mask, unified)
766 else {
767 unsafe { device.destroy_device(None) };
769 return Err(InitError::DeviceUnavailable);
770 };
771
772 let bindings = [vk::DescriptorSetLayoutBinding::default()
776 .binding(0)
777 .descriptor_type(vk::DescriptorType::STORAGE_BUFFER)
778 .descriptor_count(physical.tuning.buffers)
779 .stage_flags(vk::ShaderStageFlags::COMPUTE)];
780 let layout_info = vk::DescriptorSetLayoutCreateInfo::default().bindings(&bindings);
781 let set_layout = match unsafe { device.create_descriptor_set_layout(&layout_info, None) } {
783 Ok(layout) => layout,
784 Err(_) => {
785 unsafe { device.destroy_device(None) };
787 return Err(InitError::DeviceCreationFailed);
788 }
789 };
790 let set_layouts = [set_layout];
791 let pipeline_layout_info =
792 vk::PipelineLayoutCreateInfo::default().set_layouts(&set_layouts);
793 let pipeline_layout =
795 match unsafe { device.create_pipeline_layout(&pipeline_layout_info, None) } {
796 Ok(layout) => layout,
797 Err(_) => {
798 unsafe {
800 device.destroy_descriptor_set_layout(set_layout, None);
801 device.destroy_device(None);
802 }
803 return Err(InitError::DeviceCreationFailed);
804 }
805 };
806
807 let pipeline_cache = match unsafe {
809 device.create_pipeline_cache(&vk::PipelineCacheCreateInfo::default(), None)
810 } {
811 Ok(cache) => cache,
812 Err(_) => {
813 unsafe {
815 device.destroy_pipeline_layout(pipeline_layout, None);
816 device.destroy_descriptor_set_layout(set_layout, None);
817 device.destroy_device(None);
818 }
819 return Err(InitError::DeviceCreationFailed);
820 }
821 };
822
823 let info = device_info(&physical, memory_plan);
824 let host_memory = physical
825 .host_import_alignment
826 .map(|_| ash::ext::external_memory_host::Device::new(&instance.instance, &device));
827 Ok(Rc::new(Self {
828 device,
829 physical,
830 queue,
831 set_layout,
832 pipeline_layout,
833 pipeline_cache,
834 modules: RefCell::new(HashMap::new()),
835 memory_plan,
836 host_memory,
837 info,
838 poisoned: Cell::new(false),
839 counters: Counters::default(),
840 _instance: instance,
841 }))
842 }
843
844 fn fail(&self, result: vk::Result) -> BackendError {
846 if result == vk::Result::ERROR_DEVICE_LOST {
847 self.poisoned.set(true);
848 }
849 backend_error(result)
850 }
851
852 fn check_live(&self) -> Result<(), BackendError> {
853 if self.poisoned.get() {
854 Err(BackendError::DeviceLost)
855 } else {
856 Ok(())
857 }
858 }
859
860 fn module(&self, key: KernelKey) -> Rc<[u32]> {
862 Rc::clone(
863 self.modules
864 .borrow_mut()
865 .entry(key)
866 .or_insert_with(|| Rc::from(key.assemble())),
867 )
868 }
869
870 fn wait_idle(&self) {
873 if let Err(result) = unsafe { self.device.device_wait_idle() } {
875 self.fail(result);
876 }
877 }
878}
879
880impl Drop for Shared {
881 fn drop(&mut self) {
882 unsafe {
886 let _ = self.device.device_wait_idle();
887 self.device
888 .destroy_pipeline_cache(self.pipeline_cache, None);
889 self.device
890 .destroy_pipeline_layout(self.pipeline_layout, None);
891 self.device
892 .destroy_descriptor_set_layout(self.set_layout, None);
893 self.device.destroy_device(None);
894 }
895 }
896}
897
898fn device_info(physical: &PhysicalDeviceRecord, memory_plan: MemoryPlan) -> DeviceInfo {
899 let largest_heap = physical.memory.memory_heaps[..physical.memory.memory_heap_count as usize]
900 .iter()
901 .map(|heap| heap.size)
902 .max()
903 .unwrap_or(0);
904 let max_buffer_bytes = (u64::from(physical.limits.max_storage_buffer_range).min(largest_heap)
909 / WORD_BYTES
910 * WORD_BYTES)
911 .max(WORD_BYTES);
912 DeviceInfo {
913 identity: DeviceIdentity {
914 uuid: physical.uuid,
915 class: physical.class(),
916 vendor_id: physical.vendor_id,
917 device_id: physical.device_id,
918 },
919 capabilities: memory_plan.capabilities(),
920 limits: DeviceLimits {
921 max_contexts: MAX_CONTEXTS,
922 max_buffers_per_context: MAX_BUFFERS_PER_CONTEXT,
923 max_programs_per_context: MAX_PROGRAMS_PER_CONTEXT,
924 max_queues_per_context: MAX_QUEUES_PER_CONTEXT,
925 max_events_per_context: RING_DEPTH,
926 max_bindings_per_submission: physical.tuning.max_bindings(),
927 max_buffer_bytes,
928 max_artifact_bytes: MAX_TOSA_ARTIFACT_BYTES,
929 },
930 }
931}
932
933struct Slot {
935 command_buffer: vk::CommandBuffer,
936 fence: vk::Fence,
937 descriptor_set: vk::DescriptorSet,
938}
939
940struct ContextInner {
942 shared: Rc<Shared>,
943 id: u64,
944 command_pool: vk::CommandPool,
945 descriptor_pool: vk::DescriptorPool,
946 slots: Vec<Slot>,
947 free_slots: RefCell<Vec<u16>>,
949 transfer_command_buffer: vk::CommandBuffer,
951 transfer_fence: vk::Fence,
952}
953
954impl ContextInner {
955 fn create(shared: &Rc<Shared>, id: u64) -> Result<Rc<Self>, BackendError> {
956 let device = &shared.device;
957 let pool_info = vk::CommandPoolCreateInfo::default()
958 .flags(vk::CommandPoolCreateFlags::RESET_COMMAND_BUFFER)
959 .queue_family_index(shared.physical.queue_family);
960 let command_pool = unsafe { device.create_command_pool(&pool_info, None) }
962 .map_err(|result| shared.fail(result))?;
963 let mut partial = PartialContext {
964 shared,
965 command_pool,
966 descriptor_pool: vk::DescriptorPool::null(),
967 fences: Vec::new(),
968 };
969
970 let allocate_info = vk::CommandBufferAllocateInfo::default()
971 .command_pool(command_pool)
972 .level(vk::CommandBufferLevel::PRIMARY)
973 .command_buffer_count(RING_DEPTH + 1);
974 let command_buffers = unsafe { device.allocate_command_buffers(&allocate_info) }
976 .map_err(|result| shared.fail(result))?;
977
978 let pool_sizes = [vk::DescriptorPoolSize {
980 ty: vk::DescriptorType::STORAGE_BUFFER,
981 descriptor_count: RING_DEPTH * shared.physical.tuning.buffers,
982 }];
983 let descriptor_pool_info = vk::DescriptorPoolCreateInfo::default()
984 .max_sets(RING_DEPTH)
985 .pool_sizes(&pool_sizes);
986 partial.descriptor_pool =
988 unsafe { device.create_descriptor_pool(&descriptor_pool_info, None) }
989 .map_err(|result| shared.fail(result))?;
990 let set_layouts = vec![shared.set_layout; RING_DEPTH as usize];
991 let set_info = vk::DescriptorSetAllocateInfo::default()
992 .descriptor_pool(partial.descriptor_pool)
993 .set_layouts(&set_layouts);
994 let descriptor_sets = unsafe { device.allocate_descriptor_sets(&set_info) }
996 .map_err(|result| shared.fail(result))?;
997
998 for _ in 0..=RING_DEPTH {
999 let fence = unsafe { device.create_fence(&vk::FenceCreateInfo::default(), None) }
1001 .map_err(|result| shared.fail(result))?;
1002 partial.fences.push(fence);
1003 }
1004
1005 let (transfer_command_buffer, ring_command_buffers) = command_buffers
1006 .split_last()
1007 .expect("RING_DEPTH + 1 buffers");
1008 let transfer_fence = partial.fences.pop().expect("RING_DEPTH + 1 fences");
1009 let slots = ring_command_buffers
1010 .iter()
1011 .zip(&partial.fences)
1012 .zip(&descriptor_sets)
1013 .map(|((command_buffer, fence), descriptor_set)| Slot {
1014 command_buffer: *command_buffer,
1015 fence: *fence,
1016 descriptor_set: *descriptor_set,
1017 })
1018 .collect::<Vec<_>>();
1019 let free_slots = (0..RING_DEPTH as u16).rev().collect();
1020 let descriptor_pool = partial.descriptor_pool;
1021 partial.fences.clear();
1023 partial.descriptor_pool = vk::DescriptorPool::null();
1024 partial.command_pool = vk::CommandPool::null();
1025 increment(&shared.counters.contexts, 1);
1026 Ok(Rc::new(Self {
1027 shared: Rc::clone(shared),
1028 id,
1029 command_pool,
1030 descriptor_pool,
1031 slots,
1032 free_slots: RefCell::new(free_slots),
1033 transfer_command_buffer: *transfer_command_buffer,
1034 transfer_fence,
1035 }))
1036 }
1037
1038 fn claim_slot(&self) -> Option<u16> {
1039 self.free_slots.borrow_mut().pop()
1040 }
1041
1042 fn release_slot(&self, slot: u16) {
1043 self.free_slots.borrow_mut().push(slot);
1044 }
1045
1046 fn blocking_copy(
1048 &self,
1049 source: vk::Buffer,
1050 destination: vk::Buffer,
1051 region: vk::BufferCopy,
1052 visibility: CopyVisibility,
1053 ) -> Result<(), BackendError> {
1054 let shared = &self.shared;
1055 let device = &shared.device;
1056 let command_buffer = self.transfer_command_buffer;
1057 let fence = self.transfer_fence;
1058 let begin = vk::CommandBufferBeginInfo::default()
1059 .flags(vk::CommandBufferUsageFlags::ONE_TIME_SUBMIT);
1060 let barrier = vk::MemoryBarrier2::default()
1061 .src_stage_mask(vk::PipelineStageFlags2::COPY)
1062 .src_access_mask(vk::AccessFlags2::TRANSFER_WRITE);
1063 let barrier = match visibility {
1064 CopyVisibility::HostRead => barrier
1065 .dst_stage_mask(vk::PipelineStageFlags2::HOST)
1066 .dst_access_mask(vk::AccessFlags2::HOST_READ),
1067 CopyVisibility::Device => barrier
1068 .dst_stage_mask(
1069 vk::PipelineStageFlags2::COMPUTE_SHADER | vk::PipelineStageFlags2::COPY,
1070 )
1071 .dst_access_mask(
1072 vk::AccessFlags2::SHADER_STORAGE_READ
1073 | vk::AccessFlags2::SHADER_STORAGE_WRITE
1074 | vk::AccessFlags2::TRANSFER_READ
1075 | vk::AccessFlags2::TRANSFER_WRITE,
1076 ),
1077 };
1078 let barriers = [barrier];
1079 let dependency = vk::DependencyInfo::default().memory_barriers(&barriers);
1080 let submit_buffers =
1081 [vk::CommandBufferSubmitInfo::default().command_buffer(command_buffer)];
1082 let submits = [vk::SubmitInfo2::default().command_buffer_infos(&submit_buffers)];
1083 unsafe {
1088 device
1089 .reset_fences(&[fence])
1090 .map_err(|result| shared.fail(result))?;
1091 device
1092 .begin_command_buffer(command_buffer, &begin)
1093 .map_err(|result| shared.fail(result))?;
1094 device.cmd_copy_buffer(command_buffer, source, destination, &[region]);
1095 device.cmd_pipeline_barrier2(command_buffer, &dependency);
1096 device
1097 .end_command_buffer(command_buffer)
1098 .map_err(|result| shared.fail(result))?;
1099 device
1100 .queue_submit2(shared.queue, &submits, fence)
1101 .map_err(|result| shared.fail(result))?;
1102 match device.wait_for_fences(&[fence], true, TRANSFER_TIMEOUT_NS) {
1103 Ok(()) => Ok(()),
1104 Err(vk::Result::TIMEOUT) => {
1105 shared.poisoned.set(true);
1108 Err(BackendError::DeviceLost)
1109 }
1110 Err(result) => Err(shared.fail(result)),
1111 }
1112 }
1113 }
1114}
1115
1116#[derive(Clone, Copy, Debug, PartialEq, Eq)]
1120enum CopyVisibility {
1121 HostRead,
1123 Device,
1126}
1127
1128fn write_through_staging<S: ByteSource + ?Sized>(
1131 context: &ContextInner,
1132 destination: vk::Buffer,
1133 start: u64,
1134 data: &S,
1135 len: u64,
1136 staging: &mut Staging<'_>,
1137) -> Result<(), BackendError> {
1138 let mut done = 0_u64;
1139 let misalignment = start % WORD_BYTES;
1140 if misalignment != 0 {
1141 let prefix = (WORD_BYTES - misalignment).min(len);
1142 context.blocking_copy(
1143 destination,
1144 staging.raw.buffer,
1145 vk::BufferCopy {
1146 src_offset: start - misalignment,
1147 dst_offset: 0,
1148 size: WORD_BYTES,
1149 },
1150 CopyVisibility::HostRead,
1151 )?;
1152 let slice_start = usize::try_from(misalignment).map_err(|_| BackendError::OutOfBounds)?;
1153 let end = slice_start
1154 .checked_add(usize::try_from(prefix).map_err(|_| BackendError::OutOfBounds)?)
1155 .ok_or(BackendError::OutOfBounds)?;
1156 data.read_at(done, &mut staging.as_mut_slice()[slice_start..end])?;
1157 context.blocking_copy(
1158 staging.raw.buffer,
1159 destination,
1160 vk::BufferCopy {
1161 src_offset: 0,
1162 dst_offset: start - misalignment,
1163 size: WORD_BYTES,
1164 },
1165 CopyVisibility::Device,
1166 )?;
1167 done += prefix;
1168 }
1169 while len - done >= WORD_BYTES {
1170 let chunk = ((len - done).min(staging.bytes) / WORD_BYTES) * WORD_BYTES;
1171 let chunk_len = usize::try_from(chunk).map_err(|_| BackendError::OutOfBounds)?;
1172 data.read_at(done, &mut staging.as_mut_slice()[..chunk_len])?;
1173 context.blocking_copy(
1174 staging.raw.buffer,
1175 destination,
1176 vk::BufferCopy {
1177 src_offset: 0,
1178 dst_offset: start + done,
1179 size: chunk,
1180 },
1181 CopyVisibility::Device,
1182 )?;
1183 done += chunk;
1184 }
1185 if done < len {
1186 let tail = len - done;
1187 context.blocking_copy(
1188 destination,
1189 staging.raw.buffer,
1190 vk::BufferCopy {
1191 src_offset: start + done,
1192 dst_offset: 0,
1193 size: WORD_BYTES,
1194 },
1195 CopyVisibility::HostRead,
1196 )?;
1197 let tail_len = usize::try_from(tail).map_err(|_| BackendError::OutOfBounds)?;
1198 data.read_at(done, &mut staging.as_mut_slice()[..tail_len])?;
1199 context.blocking_copy(
1200 staging.raw.buffer,
1201 destination,
1202 vk::BufferCopy {
1203 src_offset: 0,
1204 dst_offset: start + done,
1205 size: WORD_BYTES,
1206 },
1207 CopyVisibility::Device,
1208 )?;
1209 }
1210 Ok(())
1211}
1212
1213fn read_through_staging(
1216 context: &ContextInner,
1217 source: vk::Buffer,
1218 start: u64,
1219 data: &mut dyn ByteSink,
1220 len: u64,
1221 staging: &mut Staging<'_>,
1222) -> Result<(), BackendError> {
1223 let mut done = 0_u64;
1224 let misalignment = start % WORD_BYTES;
1225 if misalignment != 0 {
1226 let prefix = (WORD_BYTES - misalignment).min(len);
1227 context.blocking_copy(
1228 source,
1229 staging.raw.buffer,
1230 vk::BufferCopy {
1231 src_offset: start - misalignment,
1232 dst_offset: 0,
1233 size: WORD_BYTES,
1234 },
1235 CopyVisibility::HostRead,
1236 )?;
1237 let slice_start = usize::try_from(misalignment).map_err(|_| BackendError::OutOfBounds)?;
1238 let end = slice_start
1239 .checked_add(usize::try_from(prefix).map_err(|_| BackendError::OutOfBounds)?)
1240 .ok_or(BackendError::OutOfBounds)?;
1241 data.write_at(done, &staging.as_mut_slice()[slice_start..end])?;
1242 done += prefix;
1243 }
1244 while len - done >= WORD_BYTES {
1245 let chunk = ((len - done).min(staging.bytes) / WORD_BYTES) * WORD_BYTES;
1246 let chunk_len = usize::try_from(chunk).map_err(|_| BackendError::OutOfBounds)?;
1247 context.blocking_copy(
1248 source,
1249 staging.raw.buffer,
1250 vk::BufferCopy {
1251 src_offset: start + done,
1252 dst_offset: 0,
1253 size: chunk,
1254 },
1255 CopyVisibility::HostRead,
1256 )?;
1257 data.write_at(done, &staging.as_mut_slice()[..chunk_len])?;
1258 done += chunk;
1259 }
1260 if done < len {
1261 let tail = usize::try_from(len - done).map_err(|_| BackendError::OutOfBounds)?;
1262 context.blocking_copy(
1263 source,
1264 staging.raw.buffer,
1265 vk::BufferCopy {
1266 src_offset: start + done,
1267 dst_offset: 0,
1268 size: WORD_BYTES,
1269 },
1270 CopyVisibility::HostRead,
1271 )?;
1272 data.write_at(done, &staging.as_mut_slice()[..tail])?;
1273 }
1274 Ok(())
1275}
1276
1277struct PartialContext<'a> {
1279 shared: &'a Shared,
1280 command_pool: vk::CommandPool,
1281 descriptor_pool: vk::DescriptorPool,
1282 fences: Vec<vk::Fence>,
1283}
1284
1285impl Drop for PartialContext<'_> {
1286 fn drop(&mut self) {
1287 let device = &self.shared.device;
1288 unsafe {
1291 for fence in self.fences.drain(..) {
1292 device.destroy_fence(fence, None);
1293 }
1294 if self.descriptor_pool != vk::DescriptorPool::null() {
1295 device.destroy_descriptor_pool(self.descriptor_pool, None);
1296 }
1297 if self.command_pool != vk::CommandPool::null() {
1298 device.destroy_command_pool(self.command_pool, None);
1299 }
1300 }
1301 }
1302}
1303
1304impl Drop for ContextInner {
1305 fn drop(&mut self) {
1306 let shared = &self.shared;
1307 if self.free_slots.borrow().len() != self.slots.len() {
1310 shared.wait_idle();
1311 }
1312 unsafe {
1315 for slot in &self.slots {
1316 shared.device.destroy_fence(slot.fence, None);
1317 }
1318 shared.device.destroy_fence(self.transfer_fence, None);
1319 shared
1320 .device
1321 .destroy_descriptor_pool(self.descriptor_pool, None);
1322 shared.device.destroy_command_pool(self.command_pool, None);
1323 }
1324 decrement(&shared.counters.contexts);
1325 }
1326}
1327
1328pub struct VulkanContext {
1330 inner: Rc<ContextInner>,
1331}
1332
1333impl std::fmt::Debug for VulkanContext {
1334 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1335 formatter
1336 .debug_struct("VulkanContext")
1337 .field("id", &self.inner.id)
1338 .finish_non_exhaustive()
1339 }
1340}
1341
1342#[derive(Default)]
1347struct BufferState {
1348 in_flight: Cell<u64>,
1349}
1350
1351pub struct VulkanBuffer {
1353 context: Rc<ContextInner>,
1354 desc: BufferDesc,
1355 buffer: vk::Buffer,
1356 memory: vk::DeviceMemory,
1357 mapped: Option<NonNull<u8>>,
1358 imported: bool,
1361 state: Rc<BufferState>,
1362}
1363
1364impl std::fmt::Debug for VulkanBuffer {
1365 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1366 formatter
1367 .debug_struct("VulkanBuffer")
1368 .field("context", &self.context.id)
1369 .field("desc", &self.desc)
1370 .field("mapped", &self.mapped.is_some())
1371 .field("imported", &self.imported)
1372 .finish_non_exhaustive()
1373 }
1374}
1375
1376impl VulkanBuffer {
1377 fn in_flight(&self) -> u64 {
1378 self.state.in_flight.get()
1379 }
1380
1381 fn mapped_at(&self, offset: usize) -> Option<*mut u8> {
1383 self.mapped
1386 .map(|pointer| unsafe { pointer.as_ptr().add(offset) })
1387 }
1388}
1389
1390impl Drop for VulkanBuffer {
1391 fn drop(&mut self) {
1392 let shared = &self.context.shared;
1393 if self.in_flight() != 0 {
1396 shared.wait_idle();
1397 }
1398 unsafe {
1403 if self.mapped.is_some() && !self.imported {
1404 shared.device.unmap_memory(self.memory);
1405 }
1406 shared.device.destroy_buffer(self.buffer, None);
1407 shared.device.free_memory(self.memory, None);
1408 }
1409 decrement(&shared.counters.buffers);
1410 }
1411}
1412
1413fn buffer_usage(physical: &PhysicalDeviceRecord) -> vk::BufferUsageFlags {
1416 let mut usage = vk::BufferUsageFlags::STORAGE_BUFFER
1417 | vk::BufferUsageFlags::TRANSFER_SRC
1418 | vk::BufferUsageFlags::TRANSFER_DST;
1419 if physical.buffer_device_address {
1420 usage |= vk::BufferUsageFlags::SHADER_DEVICE_ADDRESS;
1421 }
1422 usage
1423}
1424
1425fn probe_buffer_type_mask(
1427 device: &ash::Device,
1428 physical: &PhysicalDeviceRecord,
1429) -> Result<u32, vk::Result> {
1430 let buffer_info = vk::BufferCreateInfo::default()
1431 .size(4)
1432 .usage(buffer_usage(physical))
1433 .sharing_mode(vk::SharingMode::EXCLUSIVE);
1434 unsafe {
1436 let buffer = device.create_buffer(&buffer_info, None)?;
1437 let requirements = device.get_buffer_memory_requirements(buffer);
1438 device.destroy_buffer(buffer, None);
1439 Ok(requirements.memory_type_bits)
1440 }
1441}
1442
1443struct RawAllocation<'a> {
1445 shared: &'a Shared,
1446 buffer: vk::Buffer,
1447 memory: vk::DeviceMemory,
1448 mapped: Option<NonNull<u8>>,
1449 allocation_bytes: u64,
1450 measured_alignment: u64,
1452 memory_flags: vk::MemoryPropertyFlags,
1453}
1454
1455impl<'a> RawAllocation<'a> {
1456 fn create(
1459 shared: &'a Shared,
1460 bytes: u64,
1461 memory_type: u32,
1462 map: bool,
1463 ) -> Result<Self, BackendError> {
1464 let device = &shared.device;
1465 let rounded = bytes
1468 .checked_add(WORD_BYTES - 1)
1469 .ok_or(BackendError::ResourceLimit)?
1470 / WORD_BYTES
1471 * WORD_BYTES;
1472 let buffer_info = vk::BufferCreateInfo::default()
1473 .size(rounded)
1474 .usage(buffer_usage(&shared.physical))
1475 .sharing_mode(vk::SharingMode::EXCLUSIVE);
1476 let buffer = unsafe { device.create_buffer(&buffer_info, None) }
1478 .map_err(|result| shared.fail(result))?;
1479 let mut raw = Self {
1480 shared,
1481 buffer,
1482 memory: vk::DeviceMemory::null(),
1483 mapped: None,
1484 allocation_bytes: 0,
1485 measured_alignment: 0,
1486 memory_flags: shared.physical.memory.memory_types[memory_type as usize].property_flags,
1487 };
1488
1489 let requirements = unsafe { device.get_buffer_memory_requirements(buffer) };
1491 if requirements.memory_type_bits & (1 << memory_type) == 0 {
1492 return Err(BackendError::Incompatible);
1493 }
1494 let mut flags =
1495 vk::MemoryAllocateFlagsInfo::default().flags(vk::MemoryAllocateFlags::DEVICE_ADDRESS);
1496 let mut allocate_info = vk::MemoryAllocateInfo::default()
1497 .allocation_size(requirements.size)
1498 .memory_type_index(memory_type);
1499 if shared.physical.buffer_device_address {
1500 allocate_info = allocate_info.push_next(&mut flags);
1501 }
1502 raw.memory = unsafe { device.allocate_memory(&allocate_info, None) }
1504 .map_err(|result| shared.fail(result))?;
1505 raw.allocation_bytes = requirements.size;
1506 unsafe { device.bind_buffer_memory(buffer, raw.memory, 0) }
1508 .map_err(|result| shared.fail(result))?;
1509
1510 let mut alignment = u64::MAX;
1511 if map {
1512 let pointer = unsafe {
1515 device.map_memory(raw.memory, 0, vk::WHOLE_SIZE, vk::MemoryMapFlags::empty())
1516 }
1517 .map_err(|result| shared.fail(result))?;
1518 let pointer = NonNull::new(pointer.cast::<u8>()).ok_or(BackendError::Incompatible)?;
1519 raw.mapped = Some(pointer);
1520 alignment = alignment.min(address_alignment(pointer.as_ptr() as u64));
1521 }
1522 if shared.physical.buffer_device_address {
1523 let address_info = vk::BufferDeviceAddressInfo::default().buffer(buffer);
1524 let address = unsafe { device.get_buffer_device_address(&address_info) };
1527 alignment = alignment.min(address_alignment(address));
1528 } else if !map {
1529 alignment = alignment.min(requirements.alignment.max(1));
1531 }
1532 raw.measured_alignment = alignment;
1533 Ok(raw)
1534 }
1535
1536 unsafe fn import(
1546 shared: &'a Shared,
1547 pointer: NonNull<u8>,
1548 len: u64,
1549 ) -> Result<Self, BackendError> {
1550 let device = &shared.device;
1551 let host = shared
1552 .host_memory
1553 .as_ref()
1554 .ok_or(BackendError::Unsupported)?;
1555 let handle_type = vk::ExternalMemoryHandleTypeFlags::HOST_ALLOCATION_EXT;
1556 let mut host_properties = vk::MemoryHostPointerPropertiesEXT::default();
1557 unsafe {
1561 (host.fp().get_memory_host_pointer_properties_ext)(
1562 host.device(),
1563 handle_type,
1564 pointer.as_ptr().cast(),
1565 &mut host_properties,
1566 )
1567 }
1568 .result()
1569 .map_err(|result| shared.fail(result))?;
1570 let mut external = vk::ExternalMemoryBufferCreateInfo::default().handle_types(handle_type);
1571 let buffer_info = vk::BufferCreateInfo::default()
1572 .size(len)
1573 .usage(buffer_usage(&shared.physical))
1574 .sharing_mode(vk::SharingMode::EXCLUSIVE)
1575 .push_next(&mut external);
1576 let buffer = unsafe { device.create_buffer(&buffer_info, None) }
1578 .map_err(|result| shared.fail(result))?;
1579 let requirements = unsafe { device.get_buffer_memory_requirements(buffer) };
1581 let coherent =
1582 vk::MemoryPropertyFlags::HOST_VISIBLE | vk::MemoryPropertyFlags::HOST_COHERENT;
1583 let memory = &shared.physical.memory;
1584 let memory_type = (0..memory.memory_type_count)
1586 .filter(|&index| {
1587 host_properties.memory_type_bits & requirements.memory_type_bits & (1 << index) != 0
1588 && memory.memory_types[index as usize]
1589 .property_flags
1590 .contains(coherent)
1591 })
1592 .min_by_key(|&index| {
1593 !memory.memory_types[index as usize]
1594 .property_flags
1595 .contains(vk::MemoryPropertyFlags::DEVICE_LOCAL)
1596 });
1597 let mut raw = Self {
1598 shared,
1599 buffer,
1600 memory: vk::DeviceMemory::null(),
1601 mapped: None,
1602 allocation_bytes: 0,
1603 measured_alignment: 0,
1604 memory_flags: vk::MemoryPropertyFlags::empty(),
1605 };
1606 let memory_type = memory_type.ok_or(BackendError::Incompatible)?;
1607 raw.memory_flags = memory.memory_types[memory_type as usize].property_flags;
1608 if requirements.size > len {
1609 return Err(BackendError::ResourceLimit);
1610 }
1611 let mut import = vk::ImportMemoryHostPointerInfoEXT::default()
1612 .handle_type(handle_type)
1613 .host_pointer(pointer.as_ptr().cast());
1614 let mut flags =
1615 vk::MemoryAllocateFlagsInfo::default().flags(vk::MemoryAllocateFlags::DEVICE_ADDRESS);
1616 let mut allocate_info = vk::MemoryAllocateInfo::default()
1617 .allocation_size(len)
1618 .memory_type_index(memory_type)
1619 .push_next(&mut import);
1620 if shared.physical.buffer_device_address {
1621 allocate_info = allocate_info.push_next(&mut flags);
1622 }
1623 raw.memory = unsafe { device.allocate_memory(&allocate_info, None) }
1626 .map_err(|result| shared.fail(result))?;
1627 raw.allocation_bytes = len;
1628 unsafe { device.bind_buffer_memory(buffer, raw.memory, 0) }
1630 .map_err(|result| shared.fail(result))?;
1631 let mut alignment = address_alignment(pointer.as_ptr() as u64);
1632 if shared.physical.buffer_device_address {
1633 let address_info = vk::BufferDeviceAddressInfo::default().buffer(buffer);
1634 let address = unsafe { device.get_buffer_device_address(&address_info) };
1637 alignment = alignment.min(address_alignment(address));
1638 }
1639 raw.measured_alignment = alignment;
1640 Ok(raw)
1641 }
1642
1643 fn into_parts(self) -> (vk::Buffer, vk::DeviceMemory, Option<NonNull<u8>>) {
1645 let parts = (self.buffer, self.memory, self.mapped);
1646 std::mem::forget(self);
1647 parts
1648 }
1649}
1650
1651impl Drop for RawAllocation<'_> {
1652 fn drop(&mut self) {
1653 let device = &self.shared.device;
1654 unsafe {
1657 if self.mapped.is_some() {
1658 device.unmap_memory(self.memory);
1659 }
1660 device.destroy_buffer(self.buffer, None);
1661 if self.memory != vk::DeviceMemory::null() {
1662 device.free_memory(self.memory, None);
1663 }
1664 }
1665 }
1666}
1667
1668fn address_alignment(address: u64) -> u64 {
1670 1_u64 << address.trailing_zeros().min(40)
1671}
1672
1673struct Staging<'a> {
1675 raw: RawAllocation<'a>,
1676 bytes: u64,
1677}
1678
1679impl<'a> Staging<'a> {
1680 fn new(shared: &'a Shared, bytes: u64) -> Result<Self, BackendError> {
1681 let bytes = bytes
1682 .max(WORD_BYTES)
1683 .checked_add(WORD_BYTES - 1)
1684 .ok_or(BackendError::ResourceLimit)?
1685 / WORD_BYTES
1686 * WORD_BYTES;
1687 let raw = RawAllocation::create(shared, bytes, shared.memory_plan.host, true)?;
1688 Ok(Self { raw, bytes })
1689 }
1690
1691 fn as_mut_slice(&mut self) -> &mut [u8] {
1692 let pointer = self.raw.mapped.expect("staging is mapped");
1693 unsafe { std::slice::from_raw_parts_mut(pointer.as_ptr(), self.bytes as usize) }
1696 }
1697}
1698
1699#[derive(Default)]
1701struct ProgramState {
1702 in_flight: Cell<u32>,
1703}
1704
1705struct Dispatch {
1707 pipeline: vk::Pipeline,
1708 workgroups: [u32; 3],
1709 barrier_before: bool,
1710}
1711
1712struct Arena {
1715 buffer: vk::Buffer,
1716 memory: vk::DeviceMemory,
1717 bytes: u64,
1718 mapped: Option<NonNull<u8>>,
1721}
1722
1723pub struct VulkanProgram {
1725 context: Rc<ContextInner>,
1726 dispatches: Vec<Dispatch>,
1727 arena: Option<Arena>,
1728 plan: ProgramPlan,
1729 state: Rc<ProgramState>,
1730}
1731
1732impl std::fmt::Debug for VulkanProgram {
1733 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1734 formatter
1735 .debug_struct("VulkanProgram")
1736 .field("context", &self.context.id)
1737 .field("dispatches", &self.dispatches.len())
1738 .field("arena_bytes", &self.plan.arena_bytes)
1739 .field("slots", &self.plan.slots.len())
1740 .finish_non_exhaustive()
1741 }
1742}
1743
1744impl VulkanProgram {
1745 pub fn dispatch_count(&self) -> usize {
1747 self.dispatches.len()
1748 }
1749
1750 pub fn arena_bytes(&self) -> u64 {
1752 self.plan.arena_bytes
1753 }
1754}
1755
1756impl Drop for VulkanProgram {
1757 fn drop(&mut self) {
1758 let shared = &self.context.shared;
1759 if self.state.in_flight.get() != 0 {
1760 shared.wait_idle();
1761 }
1762 unsafe {
1766 for dispatch in &self.dispatches {
1767 shared.device.destroy_pipeline(dispatch.pipeline, None);
1768 }
1769 if let Some(arena) = &self.arena {
1770 shared.device.destroy_buffer(arena.buffer, None);
1771 shared.device.free_memory(arena.memory, None);
1772 }
1773 }
1774 decrement(&shared.counters.programs);
1775 }
1776}
1777
1778pub struct VulkanQueue {
1780 context: Rc<ContextInner>,
1781}
1782
1783impl std::fmt::Debug for VulkanQueue {
1784 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1785 formatter
1786 .debug_struct("VulkanQueue")
1787 .field("context", &self.context.id)
1788 .finish_non_exhaustive()
1789 }
1790}
1791
1792impl Drop for VulkanQueue {
1793 fn drop(&mut self) {
1794 decrement(&self.context.shared.counters.queues);
1795 }
1796}
1797
1798struct Guard {
1800 state: Rc<BufferState>,
1801 exclusive: bool,
1802}
1803
1804impl Guard {
1805 fn acquire(state: &Rc<BufferState>, exclusive: bool) -> Result<Self, BackendError> {
1806 let current = state.in_flight.get();
1807 let next = if exclusive {
1808 if current != 0 {
1809 return Err(BackendError::Busy);
1810 }
1811 EXCLUSIVE_ACCESS
1812 } else {
1813 if current >= EXCLUSIVE_ACCESS - 1 {
1814 return Err(BackendError::Busy);
1815 }
1816 current + 1
1817 };
1818 state.in_flight.set(next);
1819 Ok(Self {
1820 state: Rc::clone(state),
1821 exclusive,
1822 })
1823 }
1824}
1825
1826impl Drop for Guard {
1827 fn drop(&mut self) {
1828 let current = self.state.in_flight.get();
1829 self.state.in_flight.set(if self.exclusive {
1830 debug_assert_eq!(current, EXCLUSIVE_ACCESS);
1831 0
1832 } else {
1833 debug_assert!((1..EXCLUSIVE_ACCESS).contains(¤t));
1834 current - 1
1835 });
1836 }
1837}
1838
1839pub struct VulkanEvent {
1841 context: Rc<ContextInner>,
1842 slot: u16,
1843 program: Rc<ProgramState>,
1844 guards: RefCell<[Option<Guard>; MAX_BINDINGS_PER_SUBMISSION as usize]>,
1845 latched: Cell<Option<EventState>>,
1846 released: Cell<bool>,
1848}
1849
1850impl std::fmt::Debug for VulkanEvent {
1851 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1852 formatter
1853 .debug_struct("VulkanEvent")
1854 .field("context", &self.context.id)
1855 .field("slot", &self.slot)
1856 .field("latched", &self.latched.get())
1857 .finish_non_exhaustive()
1858 }
1859}
1860
1861impl VulkanEvent {
1862 fn latch(&self, state: EventState) -> EventState {
1865 for guard in self.guards.borrow_mut().iter_mut() {
1866 *guard = None;
1867 }
1868 let in_flight = self.program.in_flight.get();
1869 self.program.in_flight.set(in_flight.saturating_sub(1));
1870 self.latched.set(Some(state));
1871 state
1872 }
1873
1874 fn poll(&self) -> Result<EventState, BackendError> {
1876 if let Some(state) = self.latched.get() {
1877 return Ok(state);
1878 }
1879 let shared = &self.context.shared;
1880 let fence = self.context.slots[self.slot as usize].fence;
1881 match unsafe { shared.device.get_fence_status(fence) } {
1884 Ok(true) => Ok(self.latch(EventState::Complete)),
1885 Ok(false) => Ok(EventState::Pending),
1886 Err(vk::Result::ERROR_DEVICE_LOST) => {
1887 shared.poisoned.set(true);
1888 Ok(self.latch(EventState::Failed(BackendError::DeviceLost)))
1889 }
1890 Err(result) => Err(shared.fail(result)),
1891 }
1892 }
1893
1894 fn release(&self) {
1895 if self.released.replace(true) {
1896 return;
1897 }
1898 self.context.release_slot(self.slot);
1899 decrement(&self.context.shared.counters.events);
1900 }
1901}
1902
1903impl Drop for VulkanEvent {
1904 fn drop(&mut self) {
1905 if self.released.get() {
1906 return;
1907 }
1908 if self.latched.get().is_none() {
1911 let shared = &self.context.shared;
1912 let fence = self.context.slots[self.slot as usize].fence;
1913 match unsafe { shared.device.wait_for_fences(&[fence], true, u64::MAX) } {
1915 Ok(()) => {
1916 self.latch(EventState::Complete);
1917 }
1918 Err(result) => {
1919 shared.fail(result);
1920 self.latch(EventState::Failed(BackendError::DeviceLost));
1921 }
1922 }
1923 }
1924 self.release();
1925 }
1926}
1927
1928pub struct VulkanHostGate {
1937 shared: Rc<Shared>,
1938 core: Arc<GateCore>,
1939 awaited: Cell<u64>,
1941}
1942
1943struct GateCore {
1944 device: ash::Device,
1945 semaphore: vk::Semaphore,
1946 state: Mutex<GateState>,
1947}
1948
1949struct GateState {
1950 open: bool,
1952 raised: u64,
1953}
1954
1955#[derive(Clone)]
1957pub struct VulkanGateSignal {
1958 core: Arc<GateCore>,
1959}
1960
1961impl std::fmt::Debug for VulkanHostGate {
1962 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1963 formatter
1964 .debug_struct("VulkanHostGate")
1965 .field("awaited", &self.awaited.get())
1966 .finish_non_exhaustive()
1967 }
1968}
1969
1970impl std::fmt::Debug for VulkanGateSignal {
1971 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1972 formatter
1973 .debug_struct("VulkanGateSignal")
1974 .finish_non_exhaustive()
1975 }
1976}
1977
1978impl VulkanGateSignal {
1979 pub fn raise(&self, value: u64) -> Result<bool, BackendError> {
1983 let mut state = self
1985 .core
1986 .state
1987 .lock()
1988 .unwrap_or_else(|poison| poison.into_inner());
1989 if !state.open {
1990 return Ok(false);
1991 }
1992 if value > state.raised {
1993 let info = vk::SemaphoreSignalInfo::default()
1994 .semaphore(self.core.semaphore)
1995 .value(value);
1996 unsafe { self.core.device.signal_semaphore(&info) }.map_err(backend_error)?;
2000 state.raised = value;
2001 }
2002 Ok(true)
2003 }
2004}
2005
2006impl VulkanHostGate {
2007 pub fn signal(&self) -> VulkanGateSignal {
2009 VulkanGateSignal {
2010 core: Arc::clone(&self.core),
2011 }
2012 }
2013}
2014
2015impl Drop for VulkanHostGate {
2016 fn drop(&mut self) {
2017 let mut state = self
2018 .core
2019 .state
2020 .lock()
2021 .unwrap_or_else(|poison| poison.into_inner());
2022 let awaited = self.awaited.get();
2023 if awaited > state.raised {
2024 let info = vk::SemaphoreSignalInfo::default()
2025 .semaphore(self.core.semaphore)
2026 .value(awaited);
2027 if unsafe { self.core.device.signal_semaphore(&info) }.is_ok() {
2030 state.raised = awaited;
2031 }
2032 }
2033 state.open = false;
2034 drop(state);
2035 self.shared.wait_idle();
2036 unsafe {
2039 self.shared
2040 .device
2041 .destroy_semaphore(self.core.semaphore, None)
2042 };
2043 }
2044}
2045
2046pub struct VulkanAccelerator {
2047 shared: Rc<Shared>,
2048 next_id: Cell<u64>,
2049}
2050
2051impl std::fmt::Debug for VulkanAccelerator {
2052 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
2053 formatter
2054 .debug_struct("VulkanAccelerator")
2055 .field("device", &self.device_name())
2056 .finish_non_exhaustive()
2057 }
2058}
2059
2060impl VulkanAccelerator {
2061 pub fn new() -> Result<Self, InitError> {
2063 let instance = Instance::create()?;
2064 let devices = enumerate(&instance)?;
2065 let physical = devices
2066 .into_iter()
2067 .min_by_key(PhysicalDeviceRecord::rank)
2068 .ok_or(InitError::DeviceUnavailable)?;
2069 Self::open(instance, physical, VulkanOptions::default())
2070 }
2071
2072 pub fn with_device(device: &str) -> Result<Self, InitError> {
2074 Self::with_device_options(device, VulkanOptions::default())
2075 }
2076
2077 pub fn with_device_options(device: &str, options: VulkanOptions) -> Result<Self, InitError> {
2079 let instance = Instance::create()?;
2080 let physical = enumerate(&instance)?
2081 .into_iter()
2082 .find(|record| record.name == device)
2083 .ok_or(InitError::DeviceUnavailable)?;
2084 Self::open(instance, physical, options)
2085 }
2086
2087 pub fn available_devices() -> Result<Vec<String>, InitError> {
2089 let instance = Instance::create()?;
2090 Ok(enumerate(&instance)?
2091 .into_iter()
2092 .map(|record| record.name)
2093 .collect())
2094 }
2095
2096 fn open(
2097 instance: Instance,
2098 physical: PhysicalDeviceRecord,
2099 options: VulkanOptions,
2100 ) -> Result<Self, InitError> {
2101 Ok(Self {
2102 shared: Shared::open(instance, physical, options)?,
2103 next_id: Cell::new(0),
2104 })
2105 }
2106
2107 pub fn device_name(&self) -> &str {
2109 &self.shared.physical.name
2110 }
2111
2112 pub fn is_poisoned(&self) -> bool {
2114 self.shared.poisoned.get()
2115 }
2116
2117 pub fn host_gate(&self) -> Result<VulkanHostGate, BackendError> {
2119 let shared = &self.shared;
2120 shared.check_live()?;
2121 if !shared.physical.timeline_semaphore {
2122 return Err(BackendError::Unsupported);
2123 }
2124 let mut timeline = vk::SemaphoreTypeCreateInfo::default()
2125 .semaphore_type(vk::SemaphoreType::TIMELINE)
2126 .initial_value(0);
2127 let info = vk::SemaphoreCreateInfo::default().push_next(&mut timeline);
2128 let semaphore = unsafe { shared.device.create_semaphore(&info, None) }
2130 .map_err(|result| shared.fail(result))?;
2131 Ok(VulkanHostGate {
2132 shared: Rc::clone(shared),
2133 core: Arc::new(GateCore {
2134 device: shared.device.clone(),
2135 semaphore,
2136 state: Mutex::new(GateState {
2137 open: true,
2138 raised: 0,
2139 }),
2140 }),
2141 awaited: Cell::new(0),
2142 })
2143 }
2144
2145 #[allow(clippy::result_large_err)]
2154 pub fn submit_after(
2155 &self,
2156 queue: &VulkanQueue,
2157 program: &VulkanProgram,
2158 bindings: &[BindingRef<'_, VulkanBuffer>],
2159 gate: &VulkanHostGate,
2160 value: u64,
2161 ) -> Result<VulkanEvent, SubmitFailure<VulkanEvent>> {
2162 if !Rc::ptr_eq(&gate.shared, &self.shared) {
2163 return Err(SubmitFailure::Rejected(BackendError::InvalidArgument));
2164 }
2165 gate.awaited.set(gate.awaited.get().max(value));
2166 let wait = Some((gate.core.semaphore, value));
2167 self.submit_waiting(queue, program, bindings, Timeout::Infinite, wait)
2168 }
2169
2170 pub fn host_import_alignment(&self) -> Option<u64> {
2173 self.shared.physical.host_import_alignment
2174 }
2175
2176 pub unsafe fn import_host_buffer(
2194 &self,
2195 context: &VulkanContext,
2196 desc: BufferDesc,
2197 memory: NonNull<u8>,
2198 len: u64,
2199 ) -> Result<AllocatedBuffer<VulkanBuffer>, BackendError> {
2200 let shared = &self.shared;
2201 shared.info.validate_buffer_desc(desc)?;
2202 shared.check_live()?;
2203 let alignment = shared
2204 .physical
2205 .host_import_alignment
2206 .ok_or(BackendError::Unsupported)?;
2207 if !matches!(desc.domain, MemoryDomain::Host | MemoryDomain::Shared) {
2208 return Err(BackendError::Unsupported);
2209 }
2210 if memory.as_ptr() as u64 % alignment != 0 || len % alignment != 0 || desc.bytes() > len {
2211 return Err(BackendError::InvalidArgument);
2212 }
2213 if shared.counters.buffers.get()
2214 >= u64::from(MAX_BUFFERS_PER_CONTEXT) * u64::from(MAX_CONTEXTS)
2215 {
2216 return Err(BackendError::ResourceLimit);
2217 }
2218 let raw = unsafe { RawAllocation::import(shared, memory, len) }?;
2220 if raw.measured_alignment < desc.alignment() {
2221 return Err(BackendError::ResourceLimit);
2222 }
2223 let mut properties = BufferProperties::DIRECT_BINDING | BufferProperties::HOST_VISIBLE;
2224 if raw
2225 .memory_flags
2226 .contains(vk::MemoryPropertyFlags::DEVICE_LOCAL)
2227 {
2228 properties |= BufferProperties::DEVICE_LOCAL;
2229 }
2230 let info = BufferInfo::new(
2231 desc,
2232 raw.allocation_bytes,
2233 raw.measured_alignment,
2234 properties,
2235 )?;
2236 let (buffer, device_memory, _) = raw.into_parts();
2237 increment(&shared.counters.buffers, 1);
2238 Ok(AllocatedBuffer::new(
2239 VulkanBuffer {
2240 context: Rc::clone(&context.inner),
2241 desc,
2242 buffer,
2243 memory: device_memory,
2244 mapped: Some(memory),
2245 imported: true,
2246 state: Rc::new(BufferState::default()),
2247 },
2248 info,
2249 ))
2250 }
2251
2252 pub fn direct_binding_admissions(&self) -> u64 {
2254 self.shared.counters.direct_binding_admissions.get()
2255 }
2256
2257 pub fn explicit_transfer_bytes(&self) -> u64 {
2259 self.shared.counters.explicit_transfer_bytes.get()
2260 }
2261
2262 pub fn live_resources(&self) -> LiveResources {
2264 let counters = &self.shared.counters;
2265 LiveResources {
2266 contexts: counters.contexts.get(),
2267 buffers: counters.buffers.get(),
2268 programs: counters.programs.get(),
2269 queues: counters.queues.get(),
2270 events: counters.events.get(),
2271 }
2272 }
2273
2274 fn next_id(&self) -> Result<u64, BackendError> {
2275 let id = self.next_id.get();
2276 if id == u64::MAX {
2277 return Err(BackendError::ResourceLimit);
2278 }
2279 self.next_id.set(id + 1);
2280 Ok(id)
2281 }
2282
2283 fn checked_range(
2284 buffer: &VulkanBuffer,
2285 offset: u64,
2286 bytes: u64,
2287 ) -> Result<(usize, usize), BackendError> {
2288 if bytes == 0 {
2289 return Err(BackendError::InvalidArgument);
2290 }
2291 let end = offset
2292 .checked_add(bytes)
2293 .filter(|end| *end <= buffer.desc.bytes())
2294 .ok_or(BackendError::OutOfBounds)?;
2295 let start = usize::try_from(offset).map_err(|_| BackendError::OutOfBounds)?;
2296 let end = usize::try_from(end).map_err(|_| BackendError::OutOfBounds)?;
2297 Ok((start, end))
2298 }
2299
2300 fn lowering_error(error: LoweringError) -> BackendError {
2301 match error {
2302 LoweringError::Parse(_) | LoweringError::Analysis(_) => BackendError::InvalidArgument,
2303 LoweringError::UnsupportedTarget => BackendError::Incompatible,
2304 LoweringError::UnsupportedGraph
2305 | LoweringError::UnsupportedType(_)
2306 | LoweringError::UnsupportedOperator(_) => BackendError::Unsupported,
2307 LoweringError::ResourceLimit => BackendError::ResourceLimit,
2308 }
2309 }
2310
2311 fn staged_write(
2313 &self,
2314 buffer: &VulkanBuffer,
2315 start: u64,
2316 data: &dyn ByteSource,
2317 len: u64,
2318 ) -> Result<(), BackendError> {
2319 let shared = &self.shared;
2320 let mut staging = Staging::new(shared, len.min(STAGING_BYTES))?;
2321 write_through_staging(
2322 &buffer.context,
2323 buffer.buffer,
2324 start,
2325 data,
2326 len,
2327 &mut staging,
2328 )?;
2329 increment(&shared.counters.explicit_transfer_bytes, len);
2330 Ok(())
2331 }
2332
2333 fn staged_read(
2335 &self,
2336 buffer: &VulkanBuffer,
2337 start: u64,
2338 data: &mut dyn ByteSink,
2339 len: u64,
2340 ) -> Result<(), BackendError> {
2341 let shared = &self.shared;
2342 let mut staging = Staging::new(shared, len.min(STAGING_BYTES))?;
2343 read_through_staging(
2344 &buffer.context,
2345 buffer.buffer,
2346 start,
2347 data,
2348 len,
2349 &mut staging,
2350 )?;
2351 increment(&shared.counters.explicit_transfer_bytes, len);
2352 Ok(())
2353 }
2354
2355 #[allow(clippy::result_large_err)]
2359 fn submit_waiting(
2360 &self,
2361 queue: &VulkanQueue,
2362 program: &VulkanProgram,
2363 bindings: &[BindingRef<'_, VulkanBuffer>],
2364 timeout: Timeout,
2365 wait: Option<(vk::Semaphore, u64)>,
2366 ) -> Result<VulkanEvent, SubmitFailure<VulkanEvent>> {
2367 let shared = &self.shared;
2368 let reject = SubmitFailure::Rejected;
2369 shared.check_live().map_err(reject)?;
2370 if let Timeout::AfterNs(_) = timeout {
2373 return Err(reject(BackendError::DeadlineExpired));
2374 }
2375 if bindings.is_empty() || bindings.len() > shared.physical.tuning.max_bindings() as usize {
2376 return Err(reject(BackendError::ResourceLimit));
2377 }
2378 if !Rc::ptr_eq(&queue.context, &program.context) {
2379 return Err(reject(BackendError::InvalidArgument));
2380 }
2381 let context = &queue.context;
2382 let plan = &program.plan;
2383
2384 let offset_alignment = shared
2387 .physical
2388 .limits
2389 .min_storage_buffer_offset_alignment
2390 .max(WORD_BYTES);
2391 let tuning = shared.physical.tuning;
2392 let mut descriptors = vec![vk::DescriptorBufferInfo::default(); tuning.buffers as usize];
2393 let mut seen = 0_u32;
2394 for binding in bindings {
2395 if !Rc::ptr_eq(&binding.buffer.context, context) {
2396 return Err(reject(BackendError::InvalidArgument));
2397 }
2398 if !binding.buffer.desc.allows_access(binding.access) {
2399 return Err(reject(BackendError::PermissionDenied));
2400 }
2401 let (start, _) =
2402 Self::checked_range(binding.buffer, binding.range.offset, binding.range.bytes())
2403 .map_err(reject)?;
2404 let index = plan
2405 .slots
2406 .iter()
2407 .position(|slot| slot.slot == binding.slot)
2408 .ok_or(reject(BackendError::Incompatible))?;
2409 if seen & (1 << index) != 0 {
2410 return Err(reject(BackendError::InvalidArgument));
2411 }
2412 seen |= 1 << index;
2413 let slot_plan = &plan.slots[index];
2414 let expected_access = match slot_plan.role {
2415 SlotRole::Input => AccessMode::Read,
2416 SlotRole::Output => AccessMode::Write,
2417 };
2418 if binding.access != expected_access {
2419 return Err(reject(BackendError::Incompatible));
2420 }
2421 if binding.range.bytes() != slot_plan.byte_len || (start as u64) % offset_alignment != 0
2426 {
2427 return Err(reject(BackendError::Incompatible));
2428 }
2429 descriptors[index] = vk::DescriptorBufferInfo {
2430 buffer: binding.buffer.buffer,
2431 offset: start as u64,
2432 range: binding.range.bytes().div_ceil(WORD_BYTES) * WORD_BYTES,
2433 };
2434 }
2435 if bindings.len() != plan.slots.len() {
2436 return Err(reject(BackendError::Incompatible));
2437 }
2438 if let Some(arena) = &program.arena {
2441 descriptors[plan.arena_buffer_index() as usize] = vk::DescriptorBufferInfo {
2442 buffer: arena.buffer,
2443 offset: 0,
2444 range: arena.bytes,
2445 };
2446 }
2447 let filler = descriptors[0];
2448 for descriptor in &mut descriptors[plan.buffer_count() as usize..] {
2449 *descriptor = filler;
2450 }
2451 for (index, binding) in bindings.iter().enumerate() {
2456 let aliased = bindings[..index].iter().any(|prior| {
2457 Rc::ptr_eq(&prior.buffer.state, &binding.buffer.state)
2458 && (prior.access != AccessMode::Read || binding.access != AccessMode::Read)
2459 });
2460 if aliased {
2461 return Err(reject(BackendError::Incompatible));
2462 }
2463 }
2464
2465 let slot_index = context
2466 .claim_slot()
2467 .ok_or(reject(BackendError::ResourceLimit))?;
2468 let mut guards: [Option<Guard>; MAX_BINDINGS_PER_SUBMISSION as usize] =
2469 [const { None }; MAX_BINDINGS_PER_SUBMISSION as usize];
2470 for (index, binding) in bindings.iter().enumerate() {
2471 let exclusive = binding.access != AccessMode::Read;
2472 match Guard::acquire(&binding.buffer.state, exclusive) {
2475 Ok(guard) => guards[index] = Some(guard),
2476 Err(error) => {
2477 drop(guards);
2478 context.release_slot(slot_index);
2479 return Err(reject(error));
2480 }
2481 }
2482 }
2483
2484 let slot = &context.slots[slot_index as usize];
2485 match self.record_and_submit(slot, program, &descriptors, wait) {
2486 Ok(()) => {}
2487 Err(vk::Result::ERROR_DEVICE_LOST) => {
2488 shared.poisoned.set(true);
2491 program
2492 .state
2493 .in_flight
2494 .set(program.state.in_flight.get() + 1);
2495 increment(&shared.counters.events, 1);
2496 let event = VulkanEvent {
2497 context: Rc::clone(context),
2498 slot: slot_index,
2499 program: Rc::clone(&program.state),
2500 guards: RefCell::new(guards),
2501 latched: Cell::new(None),
2502 released: Cell::new(false),
2503 };
2504 event.latch(EventState::Failed(BackendError::DeviceLost));
2505 return Err(SubmitFailure::Indeterminate {
2506 error: BackendError::DeviceLost,
2507 event,
2508 });
2509 }
2510 Err(result) => {
2511 drop(guards);
2514 context.release_slot(slot_index);
2515 return Err(reject(shared.fail(result)));
2516 }
2517 }
2518 program
2519 .state
2520 .in_flight
2521 .set(program.state.in_flight.get() + 1);
2522 increment(
2523 &shared.counters.direct_binding_admissions,
2524 bindings.len() as u64,
2525 );
2526 increment(&shared.counters.events, 1);
2527 Ok(VulkanEvent {
2528 context: Rc::clone(context),
2529 slot: slot_index,
2530 program: Rc::clone(&program.state),
2531 guards: RefCell::new(guards),
2532 latched: Cell::new(None),
2533 released: Cell::new(false),
2534 })
2535 }
2536
2537 fn record_and_submit(
2540 &self,
2541 slot: &Slot,
2542 program: &VulkanProgram,
2543 descriptors: &[vk::DescriptorBufferInfo],
2544 wait: Option<(vk::Semaphore, u64)>,
2545 ) -> Result<(), vk::Result> {
2546 let shared = &self.shared;
2547 let device = &shared.device;
2548 let write = vk::WriteDescriptorSet::default()
2551 .dst_set(slot.descriptor_set)
2552 .dst_binding(0)
2553 .dst_array_element(0)
2554 .descriptor_type(vk::DescriptorType::STORAGE_BUFFER)
2555 .buffer_info(descriptors);
2556 let begin = vk::CommandBufferBeginInfo::default()
2557 .flags(vk::CommandBufferUsageFlags::ONE_TIME_SUBMIT);
2558 let compute_barrier = [vk::MemoryBarrier2::default()
2565 .src_stage_mask(vk::PipelineStageFlags2::COMPUTE_SHADER)
2566 .src_access_mask(vk::AccessFlags2::SHADER_STORAGE_WRITE)
2567 .dst_stage_mask(vk::PipelineStageFlags2::COMPUTE_SHADER)
2568 .dst_access_mask(
2569 vk::AccessFlags2::SHADER_STORAGE_READ | vk::AccessFlags2::SHADER_STORAGE_WRITE,
2570 )];
2571 let compute_dependency = vk::DependencyInfo::default().memory_barriers(&compute_barrier);
2572 let host_barrier = [vk::MemoryBarrier2::default()
2576 .src_stage_mask(vk::PipelineStageFlags2::COMPUTE_SHADER)
2577 .src_access_mask(vk::AccessFlags2::SHADER_STORAGE_WRITE)
2578 .dst_stage_mask(vk::PipelineStageFlags2::HOST | vk::PipelineStageFlags2::COPY)
2579 .dst_access_mask(
2580 vk::AccessFlags2::HOST_READ
2581 | vk::AccessFlags2::TRANSFER_READ
2582 | vk::AccessFlags2::TRANSFER_WRITE,
2583 )];
2584 let host_dependency = vk::DependencyInfo::default().memory_barriers(&host_barrier);
2585 let submit_buffers =
2586 [vk::CommandBufferSubmitInfo::default().command_buffer(slot.command_buffer)];
2587 let waits: Vec<vk::SemaphoreSubmitInfo> = wait
2591 .map(|(semaphore, value)| {
2592 vk::SemaphoreSubmitInfo::default()
2593 .semaphore(semaphore)
2594 .value(value)
2595 .stage_mask(vk::PipelineStageFlags2::COMPUTE_SHADER)
2596 })
2597 .into_iter()
2598 .collect();
2599 let submits = [vk::SubmitInfo2::default()
2600 .wait_semaphore_infos(&waits)
2601 .command_buffer_infos(&submit_buffers)];
2602 unsafe {
2608 device.update_descriptor_sets(std::slice::from_ref(&write), &[]);
2609 device.reset_fences(&[slot.fence])?;
2610 device.begin_command_buffer(slot.command_buffer, &begin)?;
2611 device.cmd_bind_descriptor_sets(
2612 slot.command_buffer,
2613 vk::PipelineBindPoint::COMPUTE,
2614 shared.pipeline_layout,
2615 0,
2616 &[slot.descriptor_set],
2617 &[],
2618 );
2619 for dispatch in &program.dispatches {
2620 if dispatch.barrier_before {
2621 device.cmd_pipeline_barrier2(slot.command_buffer, &compute_dependency);
2622 }
2623 device.cmd_bind_pipeline(
2624 slot.command_buffer,
2625 vk::PipelineBindPoint::COMPUTE,
2626 dispatch.pipeline,
2627 );
2628 device.cmd_dispatch(
2629 slot.command_buffer,
2630 dispatch.workgroups[0],
2631 dispatch.workgroups[1],
2632 dispatch.workgroups[2],
2633 );
2634 }
2635 device.cmd_pipeline_barrier2(slot.command_buffer, &host_dependency);
2636 device.end_command_buffer(slot.command_buffer)?;
2637 device.queue_submit2(shared.queue, &submits, slot.fence)
2638 }
2639 }
2640
2641 fn upload_constants(
2643 &self,
2644 context: &ContextInner,
2645 arena: &Arena,
2646 plan: &ProgramPlan,
2647 ) -> Result<(), BackendError> {
2648 let shared = &self.shared;
2649 let largest = plan
2650 .constants
2651 .iter()
2652 .map(|constant| constant.bytes.len() as u64)
2653 .max()
2654 .unwrap_or(0);
2655 if largest == 0 {
2656 return Ok(());
2657 }
2658 if let Some(mapped) = arena.mapped {
2659 for constant in &plan.constants {
2660 let offset =
2661 usize::try_from(constant.offset).map_err(|_| BackendError::OutOfBounds)?;
2662 let target = unsafe {
2667 std::slice::from_raw_parts_mut(
2668 mapped.as_ptr().add(offset),
2669 constant.bytes.len(),
2670 )
2671 };
2672 target.copy_from_slice(&constant.bytes);
2673 }
2674 return Ok(());
2675 }
2676 let mut staging = Staging::new(shared, largest.min(STAGING_BYTES))?;
2677 for constant in &plan.constants {
2678 write_through_staging(
2679 context,
2680 arena.buffer,
2681 constant.offset,
2682 constant.bytes.as_slice(),
2683 constant.bytes.len() as u64,
2684 &mut staging,
2685 )?;
2686 }
2687 Ok(())
2688 }
2689}
2690
2691struct PartialProgram<'a> {
2693 shared: &'a Shared,
2694 pipelines: Vec<vk::Pipeline>,
2695 arena: Option<Arena>,
2696}
2697
2698impl Drop for PartialProgram<'_> {
2699 fn drop(&mut self) {
2700 unsafe {
2703 for pipeline in self.pipelines.drain(..) {
2704 if pipeline != vk::Pipeline::null() {
2705 self.shared.device.destroy_pipeline(pipeline, None);
2706 }
2707 }
2708 if let Some(arena) = self.arena.take() {
2709 self.shared.device.destroy_buffer(arena.buffer, None);
2710 self.shared.device.free_memory(arena.memory, None);
2711 }
2712 }
2713 }
2714}
2715
2716impl TosaCapabilityProvider for VulkanAccelerator {
2717 fn tosa_capabilities(&self) -> &'static [CapabilityDescriptor] {
2718 &[
2722 crate::VULKAN_TOSA_FP16_CAPABILITY,
2723 crate::VULKAN_TOSA_FP8_CAPABILITY,
2724 ]
2725 }
2726}
2727
2728impl Accelerator for VulkanAccelerator {
2729 type Context = VulkanContext;
2730 type Buffer = VulkanBuffer;
2731 type Program = VulkanProgram;
2732 type Queue = VulkanQueue;
2733 type Event = VulkanEvent;
2734
2735 fn device_info(&self) -> Result<DeviceInfo, BackendError> {
2736 Ok(self.shared.info)
2737 }
2738
2739 fn create_context(&self, desc: ContextDesc) -> Result<Self::Context, BackendError> {
2740 self.shared.info.validate_context_desc(desc)?;
2741 self.shared.check_live()?;
2742 if self.shared.counters.contexts.get() >= u64::from(MAX_CONTEXTS) {
2743 return Err(BackendError::ResourceLimit);
2744 }
2745 let inner = ContextInner::create(&self.shared, self.next_id()?)?;
2746 Ok(VulkanContext { inner })
2747 }
2748
2749 fn destroy_context(&self, context: Self::Context) -> Result<(), ReleaseFailure<Self::Context>> {
2750 if Rc::strong_count(&context.inner) > 1 {
2751 return Err(ReleaseFailure::Rejected {
2752 error: BackendError::Busy,
2753 resource: context,
2754 });
2755 }
2756 Ok(())
2757 }
2758
2759 fn allocate_buffer(
2760 &self,
2761 context: &Self::Context,
2762 desc: BufferDesc,
2763 ) -> Result<AllocatedBuffer<Self::Buffer>, BackendError> {
2764 let shared = &self.shared;
2765 shared.info.validate_buffer_desc(desc)?;
2766 shared.check_live()?;
2767 if shared.counters.buffers.get()
2768 >= u64::from(MAX_BUFFERS_PER_CONTEXT) * u64::from(MAX_CONTEXTS)
2769 {
2770 return Err(BackendError::ResourceLimit);
2771 }
2772 let memory_type = shared
2773 .memory_plan
2774 .for_domain(desc.domain)
2775 .ok_or(BackendError::Unsupported)?;
2776 let map = shared
2778 .memory_plan
2779 .is_mapped(&shared.physical.memory, desc.domain);
2780 let raw = RawAllocation::create(shared, desc.bytes(), memory_type, map)?;
2781 if raw.measured_alignment < desc.alignment() {
2782 return Err(BackendError::ResourceLimit);
2783 }
2784 let mut properties = BufferProperties::DIRECT_BINDING;
2785 if raw.mapped.is_some() {
2786 properties |= BufferProperties::HOST_VISIBLE;
2787 }
2788 if raw
2789 .memory_flags
2790 .contains(vk::MemoryPropertyFlags::DEVICE_LOCAL)
2791 {
2792 properties |= BufferProperties::DEVICE_LOCAL;
2793 }
2794 let info = BufferInfo::new(
2795 desc,
2796 raw.allocation_bytes,
2797 raw.measured_alignment,
2798 properties,
2799 )?;
2800 let (buffer, memory, mapped) = raw.into_parts();
2801 increment(&shared.counters.buffers, 1);
2802 Ok(AllocatedBuffer::new(
2803 VulkanBuffer {
2804 context: Rc::clone(&context.inner),
2805 desc,
2806 buffer,
2807 memory,
2808 mapped,
2809 imported: false,
2810 state: Rc::new(BufferState::default()),
2811 },
2812 info,
2813 ))
2814 }
2815
2816 fn write_buffer(
2817 &self,
2818 buffer: &mut Self::Buffer,
2819 offset: u64,
2820 data: &dyn ByteSource,
2821 ) -> Result<(), BackendError> {
2822 if !buffer
2823 .desc
2824 .usage
2825 .contains(BufferUsage::TRANSFER_DESTINATION)
2826 {
2827 return Err(BackendError::PermissionDenied);
2828 }
2829 if buffer.in_flight() != 0 {
2830 return Err(BackendError::Busy);
2831 }
2832 self.shared.check_live()?;
2833 let (start, end) = Self::checked_range(buffer, offset, data.len())?;
2834 let len = end - start;
2835 let Some(target) = buffer.mapped_at(start) else {
2836 return self.staged_write(buffer, start as u64, data, len as u64);
2837 };
2838 let target = unsafe { std::slice::from_raw_parts_mut(target, len) };
2842 match data.as_contiguous() {
2843 Some(source) if source.len() == len => target.copy_from_slice(source),
2844 Some(_) => return Err(BackendError::InvalidArgument),
2845 None => data.read_at(0, target)?,
2846 }
2847 increment(&self.shared.counters.explicit_transfer_bytes, len as u64);
2848 Ok(())
2849 }
2850
2851 fn read_buffer(
2852 &self,
2853 buffer: &Self::Buffer,
2854 offset: u64,
2855 data: &mut dyn ByteSink,
2856 ) -> Result<(), BackendError> {
2857 if !buffer.desc.usage.contains(BufferUsage::TRANSFER_SOURCE) {
2858 return Err(BackendError::PermissionDenied);
2859 }
2860 if buffer.in_flight() != 0 {
2861 return Err(BackendError::Busy);
2862 }
2863 self.shared.check_live()?;
2864 let (start, end) = Self::checked_range(buffer, offset, data.len())?;
2865 let len = end - start;
2866 let Some(source) = buffer.mapped_at(start) else {
2867 return self.staged_read(buffer, start as u64, data, len as u64);
2868 };
2869 let source = unsafe { std::slice::from_raw_parts(source.cast_const(), len) };
2873 match data.as_contiguous_mut() {
2874 Some(target) if target.len() == len => target.copy_from_slice(source),
2875 Some(_) => return Err(BackendError::InvalidArgument),
2876 None => data.write_at(0, source)?,
2877 }
2878 increment(&self.shared.counters.explicit_transfer_bytes, len as u64);
2879 Ok(())
2880 }
2881
2882 fn free_buffer(&self, buffer: Self::Buffer) -> Result<(), ReleaseFailure<Self::Buffer>> {
2883 if buffer.in_flight() != 0 {
2884 return Err(ReleaseFailure::Rejected {
2885 error: BackendError::Busy,
2886 resource: buffer,
2887 });
2888 }
2889 Ok(())
2890 }
2891
2892 fn load_program(
2893 &self,
2894 context: &Self::Context,
2895 artifact: ArtifactRef<'_>,
2896 ) -> Result<Self::Program, BackendError> {
2897 let shared = &self.shared;
2898 if artifact.payload.len() > shared.info.limits.max_artifact_bytes {
2899 return Err(BackendError::ResourceLimit);
2900 }
2901 if artifact.resident_bytes != REQUIRED_RESIDENT_BYTES {
2902 return Err(BackendError::ResourceLimit);
2903 }
2904 shared.check_live()?;
2905 if shared.counters.programs.get()
2906 >= u64::from(MAX_PROGRAMS_PER_CONTEXT) * u64::from(MAX_CONTEXTS)
2907 {
2908 return Err(BackendError::ResourceLimit);
2909 }
2910 let mut owned = Vec::new();
2911 let bytes = match artifact.payload.as_contiguous() {
2912 Some(bytes) => bytes,
2913 None => {
2914 let len = usize::try_from(artifact.payload.len())
2915 .map_err(|_| BackendError::ResourceLimit)?;
2916 owned
2917 .try_reserve_exact(len)
2918 .map_err(|_| BackendError::OutOfMemory)?;
2919 owned.resize(len, 0);
2920 artifact.payload.read_at(0, &mut owned)?;
2921 &owned
2922 }
2923 };
2924 let plan = if artifact.format == crate::nvfp4::VULKAN_NVFP4_FORMAT {
2925 lower_nvfp4(bytes).map_err(Self::lowering_error)?
2926 } else if artifact.format == virtio_accel_tosa::ARTIFACT_FORMAT {
2927 let target = virtio_accel_tosa::Target::from_identity(artifact.target)
2928 .map_err(|_| BackendError::Incompatible)?;
2929 lower_tosa(bytes, target).map_err(Self::lowering_error)?
2930 } else {
2931 return Err(BackendError::Unsupported);
2932 };
2933 let tuning = shared.physical.tuning;
2934 let limits = &shared.physical.limits;
2935 if plan.buffer_count() > tuning.buffers
2938 || plan.slots.len() > tuning.max_bindings() as usize
2939 || plan.arena_bytes > u64::from(limits.max_storage_buffer_range)
2940 {
2941 return Err(BackendError::ResourceLimit);
2942 }
2943 let mut workgroups = Vec::with_capacity(plan.dispatches.len());
2944 for dispatch in &plan.dispatches {
2945 workgroups.push(
2946 tuning
2947 .workgroups(dispatch.work, dispatch.kernel, limits)
2948 .ok_or(BackendError::ResourceLimit)?,
2949 );
2950 }
2951
2952 let mut partial = PartialProgram {
2953 shared,
2954 pipelines: Vec::new(),
2955 arena: None,
2956 };
2957 if plan.arena_bytes != 0 {
2958 let (memory_type, map) = match shared.memory_plan.device {
2960 Some(device) => (
2961 device,
2962 shared
2963 .memory_plan
2964 .is_mapped(&shared.physical.memory, MemoryDomain::Device),
2965 ),
2966 None => (shared.memory_plan.host, true),
2967 };
2968 let raw = RawAllocation::create(shared, plan.arena_bytes, memory_type, map)?;
2969 let (buffer, memory, mapped) = raw.into_parts();
2970 partial.arena = Some(Arena {
2971 buffer,
2972 memory,
2973 bytes: plan.arena_bytes,
2974 mapped,
2975 });
2976 self.upload_constants(
2977 &context.inner,
2978 partial.arena.as_ref().expect("arena set above"),
2979 &plan,
2980 )?;
2981 }
2982
2983 let device = &shared.device;
2986 let mut modules: Vec<(KernelKey, vk::ShaderModule)> = Vec::new();
2987 let mut module_for = |key: KernelKey| -> Result<vk::ShaderModule, BackendError> {
2988 if let Some((_, module)) = modules.iter().find(|(existing, _)| *existing == key) {
2989 return Ok(*module);
2990 }
2991 let code = shared.module(key);
2992 let module_info = vk::ShaderModuleCreateInfo::default().code(&code);
2993 let module = unsafe { device.create_shader_module(&module_info, None) }
2995 .map_err(|result| shared.fail(result))?;
2996 modules.push((key, module));
2997 Ok(module)
2998 };
2999 struct ModuleGuard<'a>(&'a ash::Device, Vec<vk::ShaderModule>);
3000 impl Drop for ModuleGuard<'_> {
3001 fn drop(&mut self) {
3002 for module in self.1.drain(..) {
3003 unsafe { self.0.destroy_shader_module(module, None) };
3006 }
3007 }
3008 }
3009 let mut stage_modules = Vec::with_capacity(plan.dispatches.len());
3010 for dispatch in &plan.dispatches {
3011 let key = tuning.key(dispatch.kernel);
3012 debug_assert_eq!(dispatch.spec.len() as u32, key.spec_constant_count());
3013 match module_for(key) {
3014 Ok(module) => stage_modules.push(module),
3015 Err(error) => {
3016 drop(ModuleGuard(
3017 device,
3018 modules.into_iter().map(|(_, m)| m).collect(),
3019 ));
3020 return Err(error);
3021 }
3022 }
3023 }
3024 let module_guard = ModuleGuard(device, modules.into_iter().map(|(_, m)| m).collect());
3025 let entries: Vec<Vec<vk::SpecializationMapEntry>> = plan
3026 .dispatches
3027 .iter()
3028 .map(|dispatch| {
3029 (0..dispatch.spec.len() as u32)
3030 .map(|id| vk::SpecializationMapEntry {
3031 constant_id: id,
3032 offset: id * 4,
3033 size: 4,
3034 })
3035 .collect()
3036 })
3037 .collect();
3038 let spec_data: Vec<Vec<u8>> = plan
3039 .dispatches
3040 .iter()
3041 .map(|dispatch| {
3042 dispatch
3043 .spec
3044 .iter()
3045 .flat_map(|word| word.to_ne_bytes())
3046 .collect()
3047 })
3048 .collect();
3049 let specializations: Vec<vk::SpecializationInfo<'_>> = entries
3050 .iter()
3051 .zip(&spec_data)
3052 .map(|(entries, data)| {
3053 vk::SpecializationInfo::default()
3054 .map_entries(entries)
3055 .data(data)
3056 })
3057 .collect();
3058 let stages: Vec<vk::PipelineShaderStageCreateInfo<'_>> = stage_modules
3059 .iter()
3060 .zip(&specializations)
3061 .map(|(module, specialization)| {
3062 vk::PipelineShaderStageCreateInfo::default()
3063 .stage(vk::ShaderStageFlags::COMPUTE)
3064 .module(*module)
3065 .name(c"main")
3066 .specialization_info(specialization)
3067 })
3068 .collect();
3069 let pipeline_infos: Vec<vk::ComputePipelineCreateInfo<'_>> = stages
3070 .iter()
3071 .map(|stage| {
3072 vk::ComputePipelineCreateInfo::default()
3073 .stage(*stage)
3074 .layout(shared.pipeline_layout)
3075 })
3076 .collect();
3077 if !pipeline_infos.is_empty() {
3078 let created = unsafe {
3082 device.create_compute_pipelines(shared.pipeline_cache, &pipeline_infos, None)
3083 };
3084 match created {
3085 Ok(pipelines) => partial.pipelines = pipelines,
3086 Err((pipelines, result)) => {
3087 partial.pipelines = pipelines;
3088 drop(module_guard);
3089 return Err(shared.fail(result));
3090 }
3091 }
3092 }
3093 drop(module_guard);
3094
3095 let dispatches = partial
3096 .pipelines
3097 .drain(..)
3098 .zip(&plan.dispatches)
3099 .zip(workgroups)
3100 .map(|((pipeline, dispatch), workgroups)| Dispatch {
3101 pipeline,
3102 workgroups,
3103 barrier_before: dispatch.barrier_before,
3104 })
3105 .collect();
3106 let arena = partial.arena.take();
3107 increment(&shared.counters.programs, 1);
3108 Ok(VulkanProgram {
3109 context: Rc::clone(&context.inner),
3110 dispatches,
3111 arena,
3112 plan,
3113 state: Rc::new(ProgramState::default()),
3114 })
3115 }
3116
3117 fn unload_program(&self, program: Self::Program) -> Result<(), ReleaseFailure<Self::Program>> {
3118 if program.state.in_flight.get() != 0 {
3119 return Err(ReleaseFailure::Rejected {
3120 error: BackendError::Busy,
3121 resource: program,
3122 });
3123 }
3124 Ok(())
3125 }
3126
3127 fn create_queue(
3128 &self,
3129 context: &Self::Context,
3130 desc: QueueDesc,
3131 ) -> Result<Self::Queue, BackendError> {
3132 self.shared.info.validate_queue_desc(desc)?;
3133 self.shared.check_live()?;
3134 if self.shared.counters.queues.get()
3135 >= u64::from(MAX_QUEUES_PER_CONTEXT) * u64::from(MAX_CONTEXTS)
3136 {
3137 return Err(BackendError::ResourceLimit);
3138 }
3139 increment(&self.shared.counters.queues, 1);
3140 Ok(VulkanQueue {
3141 context: Rc::clone(&context.inner),
3142 })
3143 }
3144
3145 fn destroy_queue(&self, _queue: Self::Queue) -> Result<(), ReleaseFailure<Self::Queue>> {
3146 Ok(())
3147 }
3148
3149 fn submit(
3150 &self,
3151 queue: &Self::Queue,
3152 program: &Self::Program,
3153 bindings: &[BindingRef<'_, Self::Buffer>],
3154 timeout: Timeout,
3155 ) -> Result<Self::Event, SubmitFailure<Self::Event>> {
3156 self.submit_waiting(queue, program, bindings, timeout, None)
3157 }
3158
3159 fn poll_event(&self, event: &Self::Event) -> Result<EventState, BackendError> {
3160 event.poll()
3161 }
3162
3163 fn destroy_event(&self, event: Self::Event) -> Result<(), ReleaseFailure<Self::Event>> {
3164 match event.poll() {
3165 Ok(EventState::Pending) => Err(ReleaseFailure::Rejected {
3166 error: BackendError::Busy,
3167 resource: event,
3168 }),
3169 Ok(_) => {
3170 event.release();
3171 Ok(())
3172 }
3173 Err(error) => Err(ReleaseFailure::Rejected {
3174 error,
3175 resource: event,
3176 }),
3177 }
3178 }
3179}
3180
3181#[cfg(test)]
3182mod tests {
3183 use super::*;
3184
3185 fn amd_device_coherent_memory() -> vk::PhysicalDeviceMemoryProperties {
3189 use vk::MemoryPropertyFlags as Flags;
3190 let amd = Flags::DEVICE_COHERENT_AMD | Flags::DEVICE_UNCACHED_AMD;
3191 let host_coherent = Flags::HOST_VISIBLE | Flags::HOST_COHERENT;
3192 let layout = [
3193 (Flags::DEVICE_LOCAL, 1),
3194 (Flags::DEVICE_LOCAL, 1),
3195 (host_coherent, 0),
3196 (Flags::DEVICE_LOCAL | host_coherent, 1),
3197 (Flags::DEVICE_LOCAL | host_coherent, 1),
3198 (host_coherent | Flags::HOST_CACHED, 0),
3199 (host_coherent | Flags::HOST_CACHED, 0),
3200 (Flags::DEVICE_LOCAL | amd, 1),
3201 (host_coherent | amd, 0),
3202 (Flags::DEVICE_LOCAL | host_coherent | amd, 1),
3203 (host_coherent | Flags::HOST_CACHED | amd, 0),
3204 ];
3205 let mut memory = vk::PhysicalDeviceMemoryProperties::default();
3206 for (slot, (flags, heap)) in layout.iter().enumerate() {
3207 memory.memory_types[slot] = vk::MemoryType::default()
3208 .property_flags(*flags)
3209 .heap_index(*heap);
3210 }
3211 memory.memory_type_count =
3212 u32::try_from(layout.len()).expect("the fixture declares eleven memory types");
3213 memory.memory_heap_count = 2;
3214 memory
3215 }
3216
3217 #[test]
3223 fn never_selects_memory_that_requires_an_unrequested_feature() {
3224 let memory = amd_device_coherent_memory();
3225 let plan = MemoryPlan::select(&memory, u32::MAX, false)
3226 .expect("a host-visible coherent type is present in the fixture");
3227
3228 for (domain, selected) in [
3229 ("host", Some(plan.host)),
3230 ("device", plan.device),
3231 ("shared", plan.shared),
3232 ] {
3233 let Some(index) = selected else { continue };
3234 let flags = memory.memory_types[index as usize].property_flags;
3235 assert!(
3236 !flags.intersects(
3237 vk::MemoryPropertyFlags::DEVICE_COHERENT_AMD
3238 | vk::MemoryPropertyFlags::RDMA_CAPABLE_NV
3239 ),
3240 "{domain} domain selected memory type {index}, which requires a feature the \
3241 backend never enables: property flags {:#x}",
3242 flags.as_raw(),
3243 );
3244 }
3245 }
3246
3247 #[test]
3250 fn excluding_them_strands_no_memory_domain() {
3251 let plan = MemoryPlan::select(&amd_device_coherent_memory(), u32::MAX, false)
3252 .expect("a host-visible coherent type is present in the fixture");
3253 assert!(plan.device.is_some(), "device-local domain lost");
3254 assert!(plan.shared.is_some(), "shared domain lost");
3255 }
3256}