Skip to main content

compute_mesh/
compute_mesh.rs

1//! This example shows how to initialize an empty mesh with a Handle
2//! and a render-world only usage. That buffer is then filled by a
3//! compute shader on the GPU without transferring data back
4//! to the CPU.
5//!
6//! The `mesh_allocator` is used to get references to the relevant slabs
7//! that contain the mesh data we're interested in.
8//!
9//! This example does not remove the `GenerateMesh` component after
10//! generating the mesh.
11
12use std::ops::Not;
13
14use bevy::{
15    asset::RenderAssetUsages,
16    color::palettes::tailwind::{RED_400, SKY_400},
17    core_pipeline::schedule::camera_driver,
18    mesh::Indices,
19    platform::collections::HashSet,
20    prelude::*,
21    render::{
22        extract_component::{ExtractComponent, ExtractComponentPlugin},
23        mesh::allocator::{MeshAllocator, MeshAllocatorSettings},
24        render_resource::{
25            binding_types::{storage_buffer, uniform_buffer},
26            *,
27        },
28        renderer::{RenderContext, RenderGraph, RenderQueue},
29        Render, RenderApp, RenderStartup,
30    },
31};
32
33/// This example uses a shader source file from the assets subdirectory
34const SHADER_ASSET_PATH: &str = "shaders/compute_mesh.wesl";
35
36fn main() {
37    App::new()
38        .add_plugins((
39            DefaultPlugins,
40            ComputeShaderMeshGeneratorPlugin,
41            ExtractComponentPlugin::<GenerateMesh>::default(),
42        ))
43        .insert_resource(ClearColor(Color::BLACK))
44        .add_systems(Startup, setup)
45        .run();
46}
47
48// We need a plugin to organize all the systems and render node required for this example
49struct ComputeShaderMeshGeneratorPlugin;
50impl Plugin for ComputeShaderMeshGeneratorPlugin {
51    fn build(&self, app: &mut App) {
52        let Some(render_app) = app.get_sub_app_mut(RenderApp) else {
53            return;
54        };
55
56        render_app
57            .init_resource::<ChunksToProcess>()
58            .add_systems(RenderStartup, init_compute_pipeline)
59            .add_systems(Render, prepare_chunks)
60            .add_systems(RenderGraph, compute_mesh.before(camera_driver));
61    }
62    fn finish(&self, app: &mut App) {
63        let Some(render_app) = app.get_sub_app_mut(RenderApp) else {
64            return;
65        };
66        render_app
67            .world_mut()
68            .resource_mut::<MeshAllocatorSettings>()
69            // This allows using the mesh allocator slabs as
70            // storage buffers directly in the compute shader.
71            // Which means that we can write from our compute
72            // shader directly to the allocated mesh slabs.
73            .extra_buffer_usages = BufferUsages::STORAGE;
74    }
75}
76
77/// Holds a handle to the empty mesh that should be filled
78/// by the compute shader.
79#[derive(Component, ExtractComponent, Clone)]
80#[extract_app(RenderApp)]
81struct GenerateMesh(Handle<Mesh>);
82
83fn setup(
84    mut commands: Commands,
85    mut meshes: ResMut<Assets<Mesh>>,
86    mut materials: ResMut<Assets<StandardMaterial>>,
87) {
88    // a truly empty mesh will error if used in Mesh3d
89    // so we set up the data to be what we want the compute shader to output
90    // We're using 36 indices and 24 vertices which is directly taken from
91    // the Bevy Cuboid mesh implementation.
92    //
93    // We allocate 50 spots for each attribute here because
94    // it is *very important* that the amount of data allocated here is
95    // *bigger* than (or exactly equal to) the amount of data we intend to
96    // write from the compute shader. This amount of data defines how big
97    // the buffer we get from the mesh_allocator will be, which in turn
98    // defines how big the buffer is when we're in the compute shader.
99    //
100    // If it turns out you don't need all of the space when the compute shader
101    // is writing data, you can write NaN to the rest of the data.
102    let empty_mesh = {
103        let mut mesh = Mesh::new(
104            PrimitiveTopology::TriangleList,
105            RenderAssetUsages::RENDER_WORLD,
106        )
107        .with_inserted_attribute(Mesh::ATTRIBUTE_POSITION, vec![[0.; 3]; 50])
108        .with_inserted_attribute(Mesh::ATTRIBUTE_NORMAL, vec![[0.; 3]; 50])
109        .with_inserted_attribute(Mesh::ATTRIBUTE_UV_0, vec![[0.; 2]; 50])
110        .with_inserted_indices(Indices::U32(vec![0; 50]));
111
112        mesh.asset_usage = RenderAssetUsages::RENDER_WORLD;
113        mesh
114    };
115
116    let handle = meshes.add(empty_mesh);
117
118    // we spawn two "users" of the mesh handle,
119    // but only insert `GenerateMesh` on one of them
120    // to show that the mesh handle works as usual
121    commands.spawn((
122        GenerateMesh(handle.clone()),
123        Mesh3d(handle.clone()),
124        MeshMaterial3d(materials.add(StandardMaterial {
125            base_color: RED_400.into(),
126            ..default()
127        })),
128        Transform::from_xyz(-2.5, 1.5, 0.),
129    ));
130
131    commands.spawn((
132        Mesh3d(handle),
133        MeshMaterial3d(materials.add(StandardMaterial {
134            base_color: SKY_400.into(),
135            ..default()
136        })),
137        Transform::from_xyz(2.5, 1.5, 0.),
138    ));
139
140    // some additional scene elements.
141    // This mesh specifically is here so that we don't assume
142    // mesh_allocator offsets that would only work if we had
143    // one mesh in the scene.
144    commands.spawn((
145        Mesh3d(meshes.add(Circle::new(4.0))),
146        MeshMaterial3d(materials.add(Color::WHITE)),
147        Transform::from_rotation(Quat::from_rotation_x(-std::f32::consts::FRAC_PI_2)),
148    ));
149    commands.spawn((
150        PointLight {
151            shadow_maps_enabled: true,
152            ..default()
153        },
154        Transform::from_xyz(4.0, 8.0, 4.0),
155    ));
156    // camera
157    commands.spawn((
158        Camera3d::default(),
159        Transform::from_xyz(-2.5, 4.5, 9.0).looking_at(Vec3::ZERO, Vec3::Y),
160    ));
161}
162
163/// This is called `ChunksToProcess` because this example originated
164/// from a use case of generating chunks of landscape or voxels
165/// It only exists in the render world.
166#[derive(Resource, Default)]
167struct ChunksToProcess(Vec<AssetId<Mesh>>);
168
169/// `processed` is a `HashSet` contains the `AssetId`s that have been
170/// processed. We use that to remove `AssetId`s that have already
171/// been processed, which means each unique `GenerateMesh` will result
172/// in one compute shader mesh generation process instead of generating
173/// the mesh every frame.
174fn prepare_chunks(
175    meshes_to_generate: Query<&GenerateMesh>,
176    mut chunks: ResMut<ChunksToProcess>,
177    pipeline_cache: Res<PipelineCache>,
178    pipeline: Res<ComputePipeline>,
179    mut processed: Local<HashSet<AssetId<Mesh>>>,
180) {
181    // If the pipeline isn't ready, then meshes
182    // won't be processed. So we want to wait until
183    // the pipeline is ready before considering any mesh processed.
184    if pipeline_cache
185        .get_compute_pipeline(pipeline.pipeline)
186        .is_some()
187    {
188        // get the AssetId for each Handle<Mesh>
189        // which we'll use later to get the relevant buffers
190        // from the mesh_allocator
191        let chunk_data: Vec<AssetId<Mesh>> = meshes_to_generate
192            .iter()
193            .filter_map(|gmesh| {
194                let id = gmesh.0.id();
195                processed.contains(&id).not().then_some(id)
196            })
197            .collect();
198
199        // Cache any meshes we're going to process this frame
200        for id in &chunk_data {
201            processed.insert(*id);
202        }
203
204        chunks.0 = chunk_data;
205    }
206}
207
208#[derive(Resource)]
209struct ComputePipeline {
210    layout: BindGroupLayoutDescriptor,
211    pipeline: CachedComputePipelineId,
212}
213
214// init only happens once
215fn init_compute_pipeline(
216    mut commands: Commands,
217    asset_server: Res<AssetServer>,
218    pipeline_cache: Res<PipelineCache>,
219) {
220    let layout = BindGroupLayoutDescriptor::new(
221        "",
222        &BindGroupLayoutEntries::sequential(
223            ShaderStages::COMPUTE,
224            (
225                // offsets
226                uniform_buffer::<DataRanges>(false),
227                // vertices
228                storage_buffer::<Vec<f32>>(false),
229                // indices
230                storage_buffer::<Vec<u32>>(false),
231            ),
232        ),
233    );
234    let shader = asset_server.load(SHADER_ASSET_PATH);
235    let pipeline = pipeline_cache.queue_compute_pipeline(ComputePipelineDescriptor {
236        label: Some("Mesh generation compute shader".into()),
237        layout: vec![layout.clone()],
238        shader: shader.clone(),
239        ..default()
240    });
241    commands.insert_resource(ComputePipeline { layout, pipeline });
242}
243
244// A uniform that holds the vertex and index offsets
245// for the vertex/index mesh_allocator buffer slabs
246#[derive(ShaderType)]
247struct DataRanges {
248    vertex_start: u32,
249    vertex_end: u32,
250    index_start: u32,
251    index_end: u32,
252}
253
254fn compute_mesh(
255    mut render_context: RenderContext,
256    chunks: Res<ChunksToProcess>,
257    mesh_allocator: Res<MeshAllocator>,
258    pipeline_cache: Res<PipelineCache>,
259    pipeline: Res<ComputePipeline>,
260    render_queue: Res<RenderQueue>,
261) {
262    let Some(init_pipeline) = pipeline_cache.get_compute_pipeline(pipeline.pipeline) else {
263        return;
264    };
265
266    for mesh_id in &chunks.0 {
267        info!(?mesh_id, "processing mesh");
268
269        // the mesh_allocator holds slabs of meshes, so the buffers we get here
270        // can contain more data than just the mesh we're asking for.
271        // That's why there is a range field.
272        // You should *not* touch data in these buffers that is outside of the range.
273        let vertex_buffer_slice = mesh_allocator.mesh_vertex_slice(mesh_id).unwrap();
274        let index_buffer_slice = mesh_allocator.mesh_index_slice(mesh_id).unwrap();
275
276        let first = DataRanges {
277            // there are 8 vertex data values (pos, normal, uv) per vertex
278            // and the vertex_buffer_slice.range.start is in "vertex elements"
279            // which includes all of that data, so each index is worth 8 indices
280            // to our shader code.
281            vertex_start: vertex_buffer_slice.range.start * 8,
282            vertex_end: vertex_buffer_slice.range.end * 8,
283            // but each vertex index is a single value, so the index of the
284            // vertex indices is exactly what the value is
285            index_start: index_buffer_slice.range.start,
286            index_end: index_buffer_slice.range.end,
287        };
288
289        let mut uniforms = UniformBuffer::from(first);
290        uniforms.write_buffer(render_context.render_device(), &render_queue);
291
292        // pass in the full mesh_allocator slabs as well as the first index
293        // offsets for the vertex and index buffers
294        let bind_group = render_context.render_device().create_bind_group(
295            None,
296            &pipeline_cache.get_bind_group_layout(&pipeline.layout),
297            &BindGroupEntries::sequential((
298                &uniforms,
299                vertex_buffer_slice.buffer.as_entire_buffer_binding(),
300                index_buffer_slice.buffer.as_entire_buffer_binding(),
301            )),
302        );
303
304        let mut pass =
305            render_context
306                .command_encoder()
307                .begin_compute_pass(&ComputePassDescriptor {
308                    label: Some("Mesh generation compute pass"),
309                    ..default()
310                });
311        pass.push_debug_group("compute_mesh");
312
313        pass.set_bind_group(0, &bind_group, &[]);
314        pass.set_pipeline(init_pipeline);
315        // we only dispatch 1,1,1 workgroup here, but a real compute shader
316        // would take advantage of more and larger size workgroups
317        pass.dispatch_workgroups(1, 1, 1);
318
319        pass.pop_debug_group();
320    }
321}