use crate::backend::{ComputeCommand, ComputePipelineHandle, GpuBackend, BindGroupLayoutHandle};
use crate::bind_group::{BindGroup, BindGroupLayout};
use crate::device::Device;
use crate::shader::ShaderModule;
use anyhow::Result;
use std::sync::{Arc, Mutex};
#[derive(Clone, Default)]
pub struct ComputePipelineDesc<'a> {
pub bind_group_layouts: &'a [&'a BindGroupLayout],
}
impl<'a> std::fmt::Debug for ComputePipelineDesc<'a> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ComputePipelineDesc")
.field("bind_group_layouts_count", &self.bind_group_layouts.len())
.finish()
}
}
pub struct ComputePipeline {
backend: Arc<Mutex<Box<dyn GpuBackend>>>,
pub(crate) handle: ComputePipelineHandle,
}
impl ComputePipeline {
pub fn new(
device: &Device,
compute_shader: &ShaderModule,
desc: &ComputePipelineDesc,
) -> Result<Self> {
let mut backend = device.backend.lock().unwrap();
let layout_handles: Vec<BindGroupLayoutHandle> = desc
.bind_group_layouts
.iter()
.map(|l| l.handle)
.collect();
let handle = backend.create_compute_pipeline(
device.handle,
compute_shader.handle,
&layout_handles,
)?;
Ok(Self {
backend: Arc::clone(&device.backend),
handle,
})
}
}
impl Drop for ComputePipeline {
fn drop(&mut self) {
let mut backend = self.backend.lock().unwrap();
backend.destroy_compute_pipeline(self.handle);
}
}
pub struct ComputeEncoder {
pub(crate) commands: Vec<ComputeCommand>,
}
impl ComputeEncoder {
pub fn new() -> Self {
Self {
commands: Vec::new(),
}
}
pub fn begin_compute_pass(&mut self) -> ComputePass<'_> {
ComputePass { encoder: self }
}
pub fn finish(self) -> Vec<ComputeCommand> {
self.commands
}
pub fn dispatch(&self, device: &Device) -> Result<()> {
let mut backend = device.backend.lock().unwrap();
backend.dispatch_compute(device.handle, &self.commands)
}
}
impl Default for ComputeEncoder {
fn default() -> Self {
Self::new()
}
}
pub struct ComputePass<'a> {
encoder: &'a mut ComputeEncoder,
}
impl<'a> ComputePass<'a> {
pub fn set_pipeline(&mut self, pipeline: &ComputePipeline) {
self.encoder.commands.push(ComputeCommand::SetPipeline(pipeline.handle));
}
pub fn set_bind_group(&mut self, index: u32, bind_group: &BindGroup) {
self.encoder.commands.push(ComputeCommand::SetBindGroup {
index,
bind_group: bind_group.handle,
});
}
pub fn dispatch(&mut self, workgroups_x: u32, workgroups_y: u32, workgroups_z: u32) {
self.encoder.commands.push(ComputeCommand::Dispatch {
workgroups_x,
workgroups_y,
workgroups_z,
});
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_compute_encoder_creation() {
let encoder = ComputeEncoder::new();
assert!(encoder.commands.is_empty());
}
#[test]
fn test_compute_encoder_default() {
let encoder = ComputeEncoder::default();
assert!(encoder.commands.is_empty());
}
#[test]
fn test_dispatch_command() {
let mut encoder = ComputeEncoder::new();
{
let mut pass = encoder.begin_compute_pass();
pass.dispatch(4, 2, 1);
}
let commands = encoder.finish();
assert_eq!(commands.len(), 1);
match &commands[0] {
ComputeCommand::Dispatch { workgroups_x, workgroups_y, workgroups_z } => {
assert_eq!(*workgroups_x, 4);
assert_eq!(*workgroups_y, 2);
assert_eq!(*workgroups_z, 1);
}
_ => panic!("Expected Dispatch command"),
}
}
#[test]
fn test_multiple_dispatches() {
let mut encoder = ComputeEncoder::new();
{
let mut pass = encoder.begin_compute_pass();
pass.dispatch(1, 1, 1);
pass.dispatch(8, 8, 1);
pass.dispatch(256, 1, 1);
}
let commands = encoder.finish();
assert_eq!(commands.len(), 3);
assert!(matches!(&commands[0], ComputeCommand::Dispatch { .. }));
assert!(matches!(&commands[1], ComputeCommand::Dispatch { .. }));
assert!(matches!(&commands[2], ComputeCommand::Dispatch { .. }));
}
}