//! Memory-bandwidth probe — BabelStream-style triad. //! //! Allocates three device-local buffers (A, B, C), dispatches a //! compute shader that computes `c[i] = a[i] + s * b[i]` for `i` in //! `[0, N)`, times the dispatches, and reports observed read and //! write bandwidth from the triad: //! //! ```text //! read_bps = (2 * N * sizeof(f32)) * iterations / elapsed_seconds //! write_bps = (1 * N * sizeof(f32)) * iterations / elapsed_seconds //! ``` //! //! Two reads per element (A and B), one write (C). The shader is a //! pure streaming kernel so the achieved figure tracks the device's //! DRAM throughput, not its cache or compute path. //! //! # Timing //! //! The dispatch is timed **on the device**, with a timestamp query //! written either side of it inside the command buffer, and only the //! deltas are accumulated. Queue submission and the fence wait cost //! tens of microseconds against a dispatch of a few hundred, so timing //! the host-side loop instead charges that overhead to the memory //! system and understates bandwidth by low double-digit percent — //! measured at ~13% on an RTX 5070 Ti. //! //! A queue family that reports no valid timestamp bits falls back to //! host wall-clock timing. The result records which clock produced it, //! because the two are not comparable and a report must not imply they //! are. use std::time::Instant; use ash::vk; use naga::back::spv; use naga::front::wgsl; use naga::valid; use super::vulkan_ctx::VulkanContext; // Naga's WGSL frontend does not accept `var`, so the // triad params (which are compile-time constants anyway) are inlined // as module-scope `const`s in the shader. The N literal below must // equal ELEMENTS on the Rust side; a const_assert below guards it. const TRIAD_WGSL: &str = r" const SCALAR: f32 = 2.0; const N: u32 = 16777216u; @group(0) @binding(0) var a: array; @group(0) @binding(1) var b: array; @group(0) @binding(2) var c: array; @compute @workgroup_size(256) fn main(@builtin(global_invocation_id) gid: vec3) { let i = gid.x; if (i < N) { c[i] = a[i] + SCALAR * b[i]; } } "; /// Element count per buffer. 16 Mi f32 → 64 MiB per buffer, 192 MiB /// total. Comfortably larger than any L2/LLC so the triad is DRAM-bound. const ELEMENTS: u32 = 16 * 1024 * 1024; const _: () = assert!( ELEMENTS == 16_777_216, "ELEMENTS must equal the N literal in TRIAD_WGSL" ); const ELEMENT_BYTES: u32 = 4; const BUFFER_BYTES: vk::DeviceSize = (ELEMENTS as vk::DeviceSize) * (ELEMENT_BYTES as vk::DeviceSize); const WORKGROUP_SIZE: u32 = 256; const WARMUP_ITERATIONS: u32 = 4; const MEASURE_ITERATIONS: u32 = 32; /// Which clock produced a measurement. #[derive(Clone, Copy, Debug, PartialEq, Eq)] pub enum TimingSource { /// Timestamp queries around the dispatch. Excludes submit and /// fence-wait overhead. DeviceTimestamps, /// Host wall clock around submit-and-wait. Includes per-dispatch /// submission overhead, so the figure is a lower bound. HostWallClock, } #[derive(Clone, Copy, Debug)] pub struct BandwidthResult { pub read_bps: f64, pub write_bps: f64, pub timing: TimingSource, } #[derive(Debug)] pub enum BandwidthError { ShaderParse(String), ShaderValidate(String), ShaderEmit(String), Vulkan(vk::Result), NoDeviceLocalMemory, } impl core::fmt::Display for BandwidthError { fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result { match self { Self::ShaderParse(s) => write!(f, "WGSL parse failed: {s}"), Self::ShaderValidate(s) => write!(f, "WGSL validation failed: {s}"), Self::ShaderEmit(s) => write!(f, "SPIR-V emission failed: {s}"), Self::Vulkan(r) => write!(f, "Vulkan call failed: {r:?}"), Self::NoDeviceLocalMemory => write!(f, "no DEVICE_LOCAL memory type available"), } } } impl std::error::Error for BandwidthError {} /// Compile the triad shader and run the measurement against `ctx`. pub fn measure_bandwidth(ctx: &VulkanContext) -> Result { let spirv = compile_triad_shader()?; let device = &ctx.device; let memory_type = ctx .find_memory_type(u32::MAX, vk::MemoryPropertyFlags::DEVICE_LOCAL) .ok_or(BandwidthError::NoDeviceLocalMemory)?; // Three device-local buffers, used as storage buffers. let buffers = create_buffers(ctx, memory_type)?; // SAFETY: Vulkan handle lifetimes are managed manually below. The // teardown loop at the end of `measure_bandwidth` releases every // handle in reverse creation order. let result = unsafe { dispatch_and_time(ctx, &spirv, &buffers) }; unsafe { destroy_buffers(device, &buffers) }; let (elapsed, timing) = result?; Ok(compute_bandwidth(elapsed, MEASURE_ITERATIONS, timing)) } fn compile_triad_shader() -> Result, BandwidthError> { let module = wgsl::parse_str(TRIAD_WGSL).map_err(|e| BandwidthError::ShaderParse(format!("{e:?}")))?; let info = valid::Validator::new(valid::ValidationFlags::all(), valid::Capabilities::all()) .validate(&module) .map_err(|e| BandwidthError::ShaderValidate(format!("{e:?}")))?; let options = spv::Options::default(); spv::write_vec(&module, &info, &options, None) .map_err(|e| BandwidthError::ShaderEmit(format!("{e:?}"))) } struct TriadBuffers { buffers: [vk::Buffer; 3], memories: [vk::DeviceMemory; 3], } fn create_buffers(ctx: &VulkanContext, memory_type: u32) -> Result { let device = &ctx.device; let mut buffers = [vk::Buffer::null(); 3]; let mut memories = [vk::DeviceMemory::null(); 3]; for slot in 0..3 { let info = vk::BufferCreateInfo::default() .size(BUFFER_BYTES) .usage(vk::BufferUsageFlags::STORAGE_BUFFER) .sharing_mode(vk::SharingMode::EXCLUSIVE); let buffer = unsafe { device.create_buffer(&info, None) }.map_err(BandwidthError::Vulkan)?; let reqs = unsafe { device.get_buffer_memory_requirements(buffer) }; let alloc = vk::MemoryAllocateInfo::default() .allocation_size(reqs.size) .memory_type_index(memory_type); let memory = match unsafe { device.allocate_memory(&alloc, None) } { Ok(m) => m, Err(e) => { unsafe { device.destroy_buffer(buffer, None) }; return Err(BandwidthError::Vulkan(e)); } }; if let Err(e) = unsafe { device.bind_buffer_memory(buffer, memory, 0) } { unsafe { device.free_memory(memory, None); device.destroy_buffer(buffer, None); } return Err(BandwidthError::Vulkan(e)); } buffers[slot] = buffer; memories[slot] = memory; } Ok(TriadBuffers { buffers, memories }) } unsafe fn destroy_buffers(device: &ash::Device, b: &TriadBuffers) { for i in 0..3 { unsafe { device.destroy_buffer(b.buffers[i], None); device.free_memory(b.memories[i], None); } } } #[allow(clippy::too_many_lines)] unsafe fn dispatch_and_time( ctx: &VulkanContext, spirv: &[u32], buffers: &TriadBuffers, ) -> Result<(f64, TimingSource), BandwidthError> { let device = &ctx.device; // Descriptor set layout: three storage-buffer bindings. let bindings: [vk::DescriptorSetLayoutBinding; 3] = std::array::from_fn(|i| { #[allow(clippy::cast_possible_truncation)] let binding_idx = i as u32; vk::DescriptorSetLayoutBinding::default() .binding(binding_idx) .descriptor_type(vk::DescriptorType::STORAGE_BUFFER) .descriptor_count(1) .stage_flags(vk::ShaderStageFlags::COMPUTE) }); let dsl_info = vk::DescriptorSetLayoutCreateInfo::default().bindings(&bindings); let dsl = unsafe { device.create_descriptor_set_layout(&dsl_info, None) } .map_err(BandwidthError::Vulkan)?; // Pipeline layout: no push constants — the triad params are baked // into the shader as module-scope consts. let dsls = [dsl]; let pl_info = vk::PipelineLayoutCreateInfo::default().set_layouts(&dsls); let pipeline_layout = unsafe { device.create_pipeline_layout(&pl_info, None) }.map_err(BandwidthError::Vulkan)?; // Shader module + compute pipeline. let shader_info = vk::ShaderModuleCreateInfo::default().code(spirv); let shader = unsafe { device.create_shader_module(&shader_info, None) } .map_err(BandwidthError::Vulkan)?; let stage = vk::PipelineShaderStageCreateInfo::default() .stage(vk::ShaderStageFlags::COMPUTE) .module(shader) .name(c"main"); let pipeline_info = [vk::ComputePipelineCreateInfo::default() .stage(stage) .layout(pipeline_layout)]; let pipelines = unsafe { device.create_compute_pipelines(vk::PipelineCache::null(), &pipeline_info, None) } .map_err(|(_, r)| BandwidthError::Vulkan(r))?; let pipeline = pipelines[0]; // Descriptor pool + descriptor set, wired to our three buffers. let pool_sizes = [vk::DescriptorPoolSize::default() .ty(vk::DescriptorType::STORAGE_BUFFER) .descriptor_count(3)]; let pool_info = vk::DescriptorPoolCreateInfo::default() .max_sets(1) .pool_sizes(&pool_sizes); let descriptor_pool = unsafe { device.create_descriptor_pool(&pool_info, None) } .map_err(BandwidthError::Vulkan)?; let alloc_info = vk::DescriptorSetAllocateInfo::default() .descriptor_pool(descriptor_pool) .set_layouts(&dsls); let descriptor_set = unsafe { device.allocate_descriptor_sets(&alloc_info) }.map_err(BandwidthError::Vulkan)?[0]; let buf_infos: [vk::DescriptorBufferInfo; 3] = std::array::from_fn(|i| { vk::DescriptorBufferInfo::default() .buffer(buffers.buffers[i]) .offset(0) .range(BUFFER_BYTES) }); let writes: [vk::WriteDescriptorSet; 3] = std::array::from_fn(|i| { #[allow(clippy::cast_possible_truncation)] let binding_idx = i as u32; vk::WriteDescriptorSet::default() .dst_set(descriptor_set) .dst_binding(binding_idx) .descriptor_type(vk::DescriptorType::STORAGE_BUFFER) .buffer_info(std::slice::from_ref(&buf_infos[i])) }); unsafe { device.update_descriptor_sets(&writes, &[]) }; // Command buffer that runs one triad dispatch. let cmd_info = vk::CommandBufferAllocateInfo::default() .command_pool(ctx.command_pool) .level(vk::CommandBufferLevel::PRIMARY) .command_buffer_count(1); let cmd = unsafe { device.allocate_command_buffers(&cmd_info) }.map_err(BandwidthError::Vulkan)?[0]; // A queue that cannot write timestamps leaves only the host clock. let timing = if ctx.timestamp_valid_bits > 0 { TimingSource::DeviceTimestamps } else { TimingSource::HostWallClock }; let query_pool = if timing == TimingSource::DeviceTimestamps { let info = vk::QueryPoolCreateInfo::default() .query_type(vk::QueryType::TIMESTAMP) .query_count(2); Some(unsafe { device.create_query_pool(&info, None) }.map_err(BandwidthError::Vulkan)?) } else { None }; let begin = vk::CommandBufferBeginInfo::default(); let groups = ELEMENTS.div_ceil(WORKGROUP_SIZE); // Recorded once and resubmitted: the buffer carries no // ONE_TIME_SUBMIT flag, so re-recording per iteration bought // nothing and put another host-side cost inside the timed region. // The query-pool reset lives in the command buffer so each // submission overwrites the previous pair. unsafe { device .begin_command_buffer(cmd, &begin) .map_err(BandwidthError::Vulkan)?; if let Some(pool) = query_pool { device.cmd_reset_query_pool(cmd, pool, 0, 2); device.cmd_write_timestamp(cmd, vk::PipelineStageFlags::TOP_OF_PIPE, pool, 0); } device.cmd_bind_pipeline(cmd, vk::PipelineBindPoint::COMPUTE, pipeline); device.cmd_bind_descriptor_sets( cmd, vk::PipelineBindPoint::COMPUTE, pipeline_layout, 0, &[descriptor_set], &[], ); device.cmd_dispatch(cmd, groups, 1, 1); if let Some(pool) = query_pool { device.cmd_write_timestamp(cmd, vk::PipelineStageFlags::BOTTOM_OF_PIPE, pool, 1); } device .end_command_buffer(cmd) .map_err(BandwidthError::Vulkan)?; } let submit = || -> Result<(), BandwidthError> { let info = vk::SubmitInfo::default().command_buffers(std::slice::from_ref(&cmd)); unsafe { device .queue_submit(ctx.queue, &[info], vk::Fence::null()) .map_err(BandwidthError::Vulkan)?; device .queue_wait_idle(ctx.queue) .map_err(BandwidthError::Vulkan)?; } Ok(()) }; // Ticks are only meaningful in the low `timestamp_valid_bits`, so // mask before subtracting or a wrapped counter reads as a huge // interval. let tick_mask: u64 = if ctx.timestamp_valid_bits >= 64 { u64::MAX } else { (1_u64 << ctx.timestamp_valid_bits) - 1 }; let read_dispatch_ns = |pool: vk::QueryPool| -> Result { let mut ticks = [0_u64; 2]; unsafe { device .get_query_pool_results( pool, 0, &mut ticks, vk::QueryResultFlags::TYPE_64 | vk::QueryResultFlags::WAIT, ) .map_err(BandwidthError::Vulkan)?; } let delta = (ticks[1] & tick_mask).wrapping_sub(ticks[0] & tick_mask) & tick_mask; #[allow(clippy::cast_precision_loss)] Ok(delta as f64 * f64::from(ctx.timestamp_period_ns)) }; // Warm-up: drive the device into steady state so the measured pass // reflects thermal-stable behavior, not first-launch overhead. for _ in 0..WARMUP_ITERATIONS { submit()?; } let result = match query_pool { Some(pool) => { let mut total_ns = 0.0_f64; for _ in 0..MEASURE_ITERATIONS { submit()?; total_ns += read_dispatch_ns(pool)?; } total_ns / 1e9 } None => { let start = Instant::now(); for _ in 0..MEASURE_ITERATIONS { submit()?; } start.elapsed().as_secs_f64() } }; unsafe { if let Some(pool) = query_pool { device.destroy_query_pool(pool, None); } device.free_command_buffers(ctx.command_pool, &[cmd]); device.destroy_descriptor_pool(descriptor_pool, None); device.destroy_pipeline(pipeline, None); device.destroy_shader_module(shader, None); device.destroy_pipeline_layout(pipeline_layout, None); device.destroy_descriptor_set_layout(dsl, None); } Ok((result, timing)) } #[allow(clippy::cast_precision_loss)] fn compute_bandwidth( elapsed_seconds: f64, iterations: u32, timing: TimingSource, ) -> BandwidthResult { let bytes_per_iter_read = 2.0 * (BUFFER_BYTES as f64); let bytes_per_iter_write = BUFFER_BYTES as f64; let iters = f64::from(iterations); BandwidthResult { read_bps: bytes_per_iter_read * iters / elapsed_seconds, write_bps: bytes_per_iter_write * iters / elapsed_seconds, timing, } }