use std::sync::Arc;
use vyre_driver::BackendError;
use crate::buffer::GpuBufferHandle;
pub const INDIRECT_ARGS_BYTES: u64 = 12;
pub struct IndirectArgs {
pub buffer: Arc<wgpu::Buffer>,
pub offset: u64,
}
impl IndirectArgs {
pub fn from_handle(handle: &GpuBufferHandle, offset: u64) -> Result<Self, BackendError> {
if offset & 0b11 != 0 {
return Err(BackendError::new(format!(
"indirect dispatch offset {offset} is not 4-byte aligned. Fix: align to a u32 boundary."
)));
}
if offset
.checked_add(INDIRECT_ARGS_BYTES)
.map(|end| end > handle.byte_len())
.unwrap_or(true)
{
return Err(BackendError::new(format!(
"indirect dispatch would read past buffer end (offset={offset}, args={INDIRECT_ARGS_BYTES}, buffer byte_len={}). Fix: grow the buffer or lower the offset.",
handle.byte_len()
)));
}
if !handle.usage().contains(wgpu::BufferUsages::INDIRECT) {
return Err(BackendError::new(
"indirect dispatch requires buffer with `wgpu::BufferUsages::INDIRECT`. Fix: allocate the workgroup-count buffer with INDIRECT usage.",
));
}
Ok(Self {
buffer: handle.buffer_arc(),
offset,
})
}
}
pub fn dispatch_indirect<'a>(pass: &mut wgpu::ComputePass<'a>, args: &'a IndirectArgs) {
pass.dispatch_workgroups_indirect(&args.buffer, args.offset);
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn args_bytes_is_twelve() {
assert_eq!(INDIRECT_ARGS_BYTES, 12);
}
}