Skip to main content

mesh_shader_intro/
mesh_shader_intro.rs

1//! Mesh Shaders, at a high level, replace the classic vertex shader with a compute shader.
2//! This allows generating geometry directly on the GPU and passing those primitives directly
3//! to the fragment shader without using multiple pipelines or intermediary buffers (to pass
4//! data from a compute shader to a render pipeline).
5//!
6//! A `MeshPipeline` contains:
7//! - an optional task shader (also known as amplification shader)
8//! - a mesh shader
9//! - a fragment shader
10//!
11//! Draw calls dispatch either the task shader (if one is defined) or the mesh shader (if no task shader is defined).
12//! If a task shader runs, it can run some light processing before dictating how many mesh shaders to dispatch from the gpu.
13//!
14//! There is a task payload to pass some data between task and mesh shaders, but you should *not* use this for large amounts of data.
15//! There are hard limits typically in the low tens of kb.
16//! Adding more data to the task payload is a performance consideration, and [some documentation](https://developer.nvidia.com/blog/advanced-api-performance-mesh-shaders/#not_recommended) suggests keeping the size under 236 bytes.
17//!
18//! A mesh shader generates geometry and passes the primitives directly to the fragment shader.
19//!
20//! The fragment shader operates as usual.
21//!
22//! This is a mesh shader example that runs every frame and renders hardcoded cube mesh data at dynamic world-space coordinates defined by the mesh workgroup id.
23//! The amount of mesh shaders dispatched grows (unbounded) each frame, so will eventually exceed the platform limits and consumes more resources over time.
24//! This intentionally shows performance considerations, although the code here is illustrative and not optimized.
25//!
26//! The color of each cube is controlled by the coordinate it is rendered at combined with the task payload color.
27//! There are a couple of commented-out lines of code that render the uv or normal colors from the fragment shader, which can optionally be enabled.
28//!
29//! The cube primitives are also culled based on the time global, in a sine wave pattern.
30//!
31//! The shaders are split out into separate files intentionally, to make it easier to differentiate when specific logic is being used.
32//!
33#[cfg(feature = "free_camera")]
34use bevy::camera_controller::free_camera::{FreeCamera, FreeCameraPlugin};
35use bevy::{
36    camera::{MainPassResolutionOverride, Viewport},
37    core_pipeline::{
38        core_3d::{main_opaque_pass_3d, CORE_3D_DEPTH_FORMAT},
39        Core3d, Core3dSystems,
40    },
41    material::descriptor::{MeshPipelineDescriptor, MeshState, TaskState},
42    prelude::*,
43    render::{
44        camera::ExtractedCamera,
45        globals::{GlobalsBuffer, GlobalsUniform},
46        render_resource::{
47            binding_types::uniform_buffer, BindGroupEntries, BindGroupLayoutDescriptor,
48            BindGroupLayoutEntries, CachedRenderPipelineId, ColorTargetState, ColorWrites,
49            CompareFunction, DepthBiasState, DepthStencilState, FragmentState, PipelineCache,
50            RenderPassDescriptor, ShaderStages, StencilState, StoreOp, TextureFormat,
51        },
52        renderer::RenderContext,
53        settings::{RenderCreation, WgpuFeatures, WgpuLimits, WgpuSettings},
54        view::{
55            ExtractedView, ViewDepthStencilTexture, ViewTarget, ViewUniform, ViewUniformOffset,
56            ViewUniforms,
57        },
58        RenderApp, RenderPlugin, RenderStartup,
59    },
60};
61
62fn main() {
63    App::new()
64        .add_plugins((
65            DefaultPlugins.set(RenderPlugin {
66                render_creation: RenderCreation::Automatic(Box::new(WgpuSettings {
67                    features: WgpuFeatures::EXPERIMENTAL_MESH_SHADER
68                        | WgpuFeatures::PASSTHROUGH_SHADERS,
69                    limits: WgpuLimits::default().using_recommended_minimum_mesh_shader_values(),
70                    ..default()
71                })),
72                ..default()
73            }),
74            MeshShaderDemoPlugin,
75            #[cfg(feature = "free_camera")]
76            FreeCameraPlugin,
77        ))
78        .add_systems(Startup, setup)
79        .run();
80}
81
82fn setup(
83    mut commands: Commands,
84    mut meshes: ResMut<Assets<Mesh>>,
85    mut materials: ResMut<Assets<StandardMaterial>>,
86) {
87    commands.spawn((
88        Mesh3d(meshes.add(Cuboid::new(0.5, 0.5, 0.5))),
89        MeshMaterial3d(materials.add(Color::srgb_u8(124, 144, 255))),
90        Transform::from_xyz(0.0, 0.5, 0.0),
91    ));
92    // light
93    commands.spawn((
94        PointLight {
95            shadow_maps_enabled: true,
96            ..default()
97        },
98        Transform::from_xyz(4.0, 8.0, 4.0),
99    ));
100    // camera
101    commands.spawn((
102        Camera3d::default(),
103        Transform::from_xyz(-9.0, 1.0, -9.0).looking_at(Vec3::new(9., 4., 9.), Vec3::Y),
104        // disable msaa for simplicity
105        Msaa::Off,
106        #[cfg(feature = "free_camera")]
107        FreeCamera::default(),
108    ));
109}
110
111struct MeshShaderDemoPlugin;
112impl Plugin for MeshShaderDemoPlugin {
113    fn build(&self, app: &mut App) {
114        let Some(render_app) = app.get_sub_app_mut(RenderApp) else {
115            return;
116        };
117
118        render_app
119            .add_systems(RenderStartup, init_mesh_pipelines)
120            .add_systems(
121                Core3d,
122                draw_mesh_shader_cubes
123                    .after(main_opaque_pass_3d)
124                    .in_set(Core3dSystems::MainPass),
125            );
126    }
127}
128
129fn init_mesh_pipelines(
130    mut commands: Commands,
131    asset_server: Res<AssetServer>,
132    pipeline_cache: Res<PipelineCache>,
133) {
134    let task_shader = asset_server.load::<Shader>("shaders/mesh_shader_intro/task.wesl");
135    let mesh_shader = asset_server.load::<Shader>("shaders/mesh_shader_intro/mesh.wesl");
136    let fragment_shader = asset_server.load::<Shader>("shaders/mesh_shader_intro/fragment.wesl");
137
138    let layout = BindGroupLayoutDescriptor::new(
139        "custom_mesh_shader_bind_group_layout",
140        &BindGroupLayoutEntries::sequential(
141            ShaderStages::MESH | ShaderStages::TASK | ShaderStages::FRAGMENT,
142            (
143                uniform_buffer::<GlobalsUniform>(false),
144                uniform_buffer::<ViewUniform>(true),
145            ),
146        ),
147    );
148
149    let depth_stencil = DepthStencilState {
150        format: CORE_3D_DEPTH_FORMAT,
151        depth_write_enabled: Some(true),
152        depth_compare: Some(CompareFunction::GreaterEqual),
153        stencil: StencilState::default(),
154        bias: DepthBiasState::default(),
155    };
156
157    let cached_pipeline_id = pipeline_cache.queue_mesh_pipeline(MeshPipelineDescriptor {
158        label: Some("custom_mesh_shader_pipeline".into()),
159        layout: vec![layout.clone()],
160        immediate_size: 0,
161        task: Some(TaskState {
162            shader: task_shader,
163            entry_point: Some("task".into()),
164            ..default()
165        }),
166        mesh: MeshState {
167            shader: mesh_shader,
168            entry_point: Some("mesh".into()),
169            ..default()
170        },
171        primitive: Default::default(),
172        depth_stencil: Some(depth_stencil),
173        multisample: Default::default(),
174        fragment: Some(FragmentState {
175            shader: fragment_shader,
176            entry_point: Some("fragment".into()),
177            targets: vec![Some(ColorTargetState {
178                format: TextureFormat::Rgba8UnormSrgb,
179                blend: None,
180                write_mask: ColorWrites::ALL,
181            })],
182            ..default()
183        }),
184        zero_initialize_workgroup_memory: false,
185    });
186
187    commands.insert_resource(MyMeshShaderDrawNode {
188        mesh_pipeline: cached_pipeline_id,
189        layout,
190    });
191}
192
193#[derive(Resource)]
194struct MyMeshShaderDrawNode {
195    mesh_pipeline: CachedRenderPipelineId,
196    layout: BindGroupLayoutDescriptor,
197}
198
199/// The underlying `create_mesh_pipeline` returns a `RenderPipeline`, which means
200/// mesh shaders can re-use the `RenderPass` infrastructure from other examples to
201/// start a `TrackedRenderPass`.
202fn draw_mesh_shader_cubes(
203    mut views: Query<(
204        &ExtractedCamera,
205        &ExtractedView,
206        &ViewTarget,
207        &ViewDepthStencilTexture,
208        &ViewUniformOffset,
209        Option<&MainPassResolutionOverride>,
210    )>,
211    mut render_context: RenderContext,
212    data: Res<MyMeshShaderDrawNode>,
213    view_uniforms: Res<ViewUniforms>,
214    globals: Res<GlobalsBuffer>,
215    pipeline_cache: Res<PipelineCache>,
216) {
217    let Some(mesh_pipeline) = pipeline_cache.get_render_pipeline(data.mesh_pipeline) else {
218        return;
219    };
220
221    for (camera, _, target, depth, view_uniform_offset, resolution_override) in &mut views {
222        let Some(view_binding) = view_uniforms.uniforms.binding() else {
223            return;
224        };
225        let Some(globals_binding) = globals.buffer.binding() else {
226            return;
227        };
228        let bind_group = render_context.render_device().create_bind_group(
229            "custom_task_mesh_bind_group",
230            &pipeline_cache.get_bind_group_layout(&data.layout),
231            &BindGroupEntries::sequential((globals_binding, view_binding)),
232        );
233
234        {
235            let mut pass = render_context.begin_tracked_render_pass(RenderPassDescriptor {
236                label: Some("custom_mesh_shader_pass"),
237                // Write directly to the view target
238                color_attachments: &[Some(target.get_color_attachment())],
239                depth_stencil_attachment: Some(depth.get_attachment(StoreOp::Store)),
240                timestamp_writes: None,
241                occlusion_query_set: None,
242                multiview_mask: None,
243            });
244
245            pass.set_render_pipeline(mesh_pipeline);
246            pass.set_bind_group(0, &bind_group, &[view_uniform_offset.offset]);
247            if let Some(viewport) =
248                Viewport::from_viewport_and_override(camera.viewport.as_ref(), resolution_override)
249            {
250                pass.set_camera_viewport(&viewport);
251            }
252
253            // Since this MeshPipeline has a task shader, this call
254            // dispatches the task shader workgroup
255            pass.draw_mesh_tasks(1, 1, 1);
256        }
257    }
258}