From 93b86f23b005ff4789337bff5cb0861053d058a1 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?charlotte=20=F0=9F=8C=B8?= Date: Mon, 28 Sep 2026 13:21:23 -0700 Subject: [PATCH] Compute perf. --- crates/processing_render/src/compute.rs | 37 +++++++-- crates/processing_render/src/lib.rs | 82 ++++++++++++++----- .../processing_render/src/particles/grid.rs | 10 +-- .../processing_render/src/particles/scan.rs | 8 +- 4 files changed, 99 insertions(+), 38 deletions(-) diff --git a/crates/processing_render/src/compute.rs b/crates/processing_render/src/compute.rs index 4b56fe6a..9e593ed9 100644 --- a/crates/processing_render/src/compute.rs +++ b/crates/processing_render/src/compute.rs @@ -112,17 +112,40 @@ pub fn create_buffer_with_data( .id() } -pub fn write_buffer_cpu( - In((handle, offset, data)): In<(Handle, u64, Vec)>, +/// Writes to the GPU buffer, false if it isn't prepared yet. +pub fn write_buffer_gpu( + InRef((handle, offset, data)): InRef<(Handle, u64, Vec)>, + gpu_buffers: Res>, + render_queue: Res, +) -> bool { + let Some(gpu_buffer) = gpu_buffers.get(handle) else { + return false; + }; + render_queue.write_buffer(&gpu_buffer.buffer, *offset, data); + true +} + +/// Writes to the buffer's asset. +pub fn write_buffer_asset( + In(((handle, offset, data), tracked)): In<((Handle, u64, Vec), bool)>, mut buffers: ResMut>, ) -> Result<()> { - let mut asset = buffers - .get_mut(&handle) - .ok_or(ProcessingError::BufferNotFound)?; + let patched = if tracked { + buffers + .get_mut(&handle) + .map(|mut asset| patch_asset(&mut asset, offset, &data)) + } else { + buffers + .get_mut_untracked(handle.id()) + .map(|asset| patch_asset(asset, offset, &data)) + }; + patched.ok_or(ProcessingError::BufferNotFound)? +} + +fn patch_asset(asset: &mut ShaderBuffer, offset: u64, data: &[u8]) -> Result<()> { let dst = asset.data.as_mut().ok_or(ProcessingError::BufferNotFound)?; let start = offset as usize; - let end = start + data.len(); - dst[start..end].copy_from_slice(&data); + dst[start..start + data.len()].copy_from_slice(data); Ok(()) } diff --git a/crates/processing_render/src/lib.rs b/crates/processing_render/src/lib.rs index 5b0cc9b2..67280944 100644 --- a/crates/processing_render/src/lib.rs +++ b/crates/processing_render/src/lib.rs @@ -2432,9 +2432,23 @@ fn buffer_write_range( data.len() ))); } - ensure_buffer_synced(app, entity)?; + let write = (handle, offset, data); + let on_gpu = app + .sub_app_mut(bevy::render::RenderApp) + .world_mut() + .run_system_cached_with(compute::write_buffer_gpu, &write) + .unwrap(); + let synced = app + .world() + .get::(entity) + .ok_or(error::ProcessingError::BufferNotFound)? + .synced; + // next read will refresh + if on_gpu && !synced { + return Ok(()); + } app.world_mut() - .run_system_cached_with(compute::write_buffer_cpu, (handle, offset, data)) + .run_system_cached_with(compute::write_buffer_asset, (write, !on_gpu)) .unwrap() }) } @@ -2527,30 +2541,54 @@ pub fn compute_set( pub fn compute_dispatch(entity: Entity, x: u32, y: u32, z: u32) -> error::Result<()> { app_mut(|app| { app.update(); + dispatch_inner(app, entity, x, y, z) + }) +} - let args = { - let world = app.world(); - let c = world - .get::(entity) - .ok_or(error::ProcessingError::ComputeNotFound)?; - let mesh_bindings = compute::resolve_mesh_bindings(world, c)?; - ( - c.pipeline_id, - c.bind_group_layout_descriptors.clone(), - c.shader.clone(), - mesh_bindings, - x, - y, - z, - ) - }; - app.sub_app_mut(bevy::render::RenderApp) - .world_mut() - .run_system_cached_with(compute::dispatch, args) - .unwrap() +/// Dispatch without an `app.update()` first; only updates if the pipeline isn't compiled yet. +pub(crate) fn compute_dispatch_no_update( + entity: Entity, + x: u32, + y: u32, + z: u32, +) -> error::Result<()> { + app_mut(|app| match dispatch_inner(app, entity, x, y, z) { + Err(error::ProcessingError::PipelineNotReady(_)) => { + app.update(); + dispatch_inner(app, entity, x, y, z) + } + r => r, }) } +fn dispatch_inner(app: &mut App, entity: Entity, x: u32, y: u32, z: u32) -> error::Result<()> { + let args = { + let world = app.world(); + let c = world + .get::(entity) + .ok_or(error::ProcessingError::ComputeNotFound)?; + let mesh_bindings = compute::resolve_mesh_bindings(world, c)?; + ( + c.pipeline_id, + c.bind_group_layout_descriptors.clone(), + c.shader.clone(), + mesh_bindings, + x, + y, + z, + ) + }; + app.sub_app_mut(bevy::render::RenderApp) + .world_mut() + .run_system_cached_with(compute::dispatch, args) + .unwrap()?; + // The kernel may have written its read-write buffers, so their assets are no longer synced. + app.world_mut() + .run_system_cached(compute::invalidate_rw_buffers) + .unwrap(); + Ok(()) +} + pub fn compute_destroy(entity: Entity) -> error::Result<()> { app_mut(|app| { app.world_mut() diff --git a/crates/processing_render/src/particles/grid.rs b/crates/processing_render/src/particles/grid.rs index 0a31f069..0e778111 100644 --- a/crates/processing_render/src/particles/grid.rs +++ b/crates/processing_render/src/particles/grid.rs @@ -6,7 +6,7 @@ use processing_core::error::Result; use crate::particles::scan::prefix_sum_u32; use crate::shader_value::ShaderValue; -use crate::{buffer_create, compute_create, compute_dispatch, compute_set, shader_load}; +use crate::{buffer_create, compute_create, compute_dispatch_no_update, compute_set, shader_load}; const CLEAR_SHADER: &str = "embedded://processing_render/particles/kernels/grid_clear.wgsl"; const COUNT_SHADER: &str = "embedded://processing_render/particles/kernels/grid_count.wgsl"; @@ -89,24 +89,24 @@ pub fn grid_build(grid: &Grid, position: Entity) -> Result<()> { let num_cells = grid.params.num_cells(); compute_set(clear, "counts", ShaderValue::Buffer(grid.offsets))?; - compute_dispatch(clear, (num_cells + 1).div_ceil(CLEAR_WG), 1, 1)?; + compute_dispatch_no_update(clear, (num_cells + 1).div_ceil(CLEAR_WG), 1, 1)?; compute_set(count, "position", ShaderValue::Buffer(position))?; compute_set(count, "counts", ShaderValue::Buffer(grid.offsets))?; set_domain(count, &grid.params)?; - compute_dispatch(count, grid.capacity.div_ceil(PARTICLE_WG), 1, 1)?; + compute_dispatch_no_update(count, grid.capacity.div_ceil(PARTICLE_WG), 1, 1)?; prefix_sum_u32(grid.offsets)?; compute_set(copy, "starts", ShaderValue::Buffer(grid.offsets))?; compute_set(copy, "cursor", ShaderValue::Buffer(grid.cursor))?; - compute_dispatch(copy, num_cells.div_ceil(COPY_WG), 1, 1)?; + compute_dispatch_no_update(copy, num_cells.div_ceil(COPY_WG), 1, 1)?; compute_set(scatter, "position", ShaderValue::Buffer(position))?; compute_set(scatter, "cursor", ShaderValue::Buffer(grid.cursor))?; compute_set(scatter, "sorted", ShaderValue::Buffer(grid.sorted))?; set_domain(scatter, &grid.params)?; - compute_dispatch(scatter, grid.capacity.div_ceil(PARTICLE_WG), 1, 1)?; + compute_dispatch_no_update(scatter, grid.capacity.div_ceil(PARTICLE_WG), 1, 1)?; Ok(()) } diff --git a/crates/processing_render/src/particles/scan.rs b/crates/processing_render/src/particles/scan.rs index bab519e7..e433f683 100644 --- a/crates/processing_render/src/particles/scan.rs +++ b/crates/processing_render/src/particles/scan.rs @@ -6,8 +6,8 @@ use processing_core::error::Result; use crate::shader_value::ShaderValue; use crate::{ - buffer_create, buffer_destroy, buffer_size, compute_create, compute_dispatch, compute_set, - shader_load, + buffer_create, buffer_destroy, buffer_size, compute_create, compute_dispatch_no_update, + compute_set, shader_load, }; const BLOCK: u64 = 256; @@ -67,7 +67,7 @@ pub fn prefix_sum_u32(buffer: Entity) -> Result<()> { compute_set(block, "data", ShaderValue::Buffer(level_bufs[lvl]))?; compute_set(block, "block_sums", ShaderValue::Buffer(sums))?; - compute_dispatch(block, num_blocks as u32, 1, 1)?; + compute_dispatch_no_update(block, num_blocks as u32, 1, 1)?; level_bufs.push(sums); level_ns.push(num_blocks); @@ -82,7 +82,7 @@ pub fn prefix_sum_u32(buffer: Entity) -> Result<()> { let num_blocks = level_ns[k].div_ceil(BLOCK).max(1); compute_set(add, "data", ShaderValue::Buffer(level_bufs[k]))?; compute_set(add, "block_sums", ShaderValue::Buffer(level_bufs[k + 1]))?; - compute_dispatch(add, num_blocks as u32, 1, 1)?; + compute_dispatch_no_update(add, num_blocks as u32, 1, 1)?; } Ok(())