Skip to main content

custom_post_processing/
custom_post_processing.rs

1//! This example shows how to create a custom post-processing effect that runs after the main pass
2//! and reads the texture generated by the main pass.
3//!
4//! The example shader is a very simple implementation of chromatic aberration.
5//! To adapt this example for 2D, replace all instances of 3D structures (such as `Core3d`, etc.) with their corresponding 2D counterparts.
6//!
7//! This is a fairly low level example and assumes some familiarity with rendering concepts and wgpu.
8
9use bevy::{
10    core_pipeline::{schedule::Core3d, Core3dSystems, FullscreenShader},
11    prelude::*,
12    render::{
13        extract_component::{
14            ComponentUniforms, DynamicUniformIndex, ExtractComponent, ExtractComponentPlugin,
15            UniformComponentPlugin,
16        },
17        render_resource::{
18            binding_types::{sampler, texture_2d, uniform_buffer},
19            *,
20        },
21        renderer::{RenderContext, RenderDevice, ViewQuery},
22        view::ViewTarget,
23        Render, RenderApp, RenderStartup, RenderSystems,
24    },
25};
26
27/// This example uses a shader source file from the assets subdirectory
28const SHADER_ASSET_PATH: &str = "shaders/post_processing.wesl";
29
30fn main() {
31    App::new()
32        .add_plugins((DefaultPlugins, PostProcessPlugin))
33        .add_systems(Startup, setup)
34        .add_systems(Update, (rotate, update_settings))
35        .run();
36}
37
38/// It is generally encouraged to set up post processing effects as a plugin
39struct PostProcessPlugin;
40
41impl Plugin for PostProcessPlugin {
42    fn build(&self, app: &mut App) {
43        app.add_plugins((
44            // The settings will be a component that lives in the main world but will
45            // be extracted to the render world every frame.
46            // This makes it possible to control the effect from the main world.
47            // This plugin will take care of extracting it automatically.
48            // It's important to derive [`ExtractComponent`] on [`PostProcessSettings`]
49            // for this plugin to work correctly.
50            ExtractComponentPlugin::<PostProcessSettings>::default(),
51            // The settings will also be the data used in the shader.
52            // This plugin will prepare the component for the GPU by creating a uniform buffer
53            // and writing the data to that buffer every frame.
54            UniformComponentPlugin::<PostProcessSettings>::default(),
55        ));
56
57        // We need to get the render app from the main app
58        let Some(render_app) = app.get_sub_app_mut(RenderApp) else {
59            return;
60        };
61
62        render_app
63            .add_systems(RenderStartup, init_post_process_pipeline)
64            .add_systems(
65                Render,
66                prepare_bind_groups.in_set(RenderSystems::PrepareBindGroups),
67            )
68            .add_systems(
69                Core3d,
70                post_process_system.in_set(Core3dSystems::PostProcess),
71            );
72    }
73}
74
75/// Holds the bind groups for both main textures
76///
77/// We can't know ahead of time which one is the source or destination so we create a bind group
78/// for both
79#[derive(Component)]
80struct PostProcessBindGroups {
81    a: (TextureViewId, BindGroup),
82    b: (TextureViewId, BindGroup),
83}
84
85/// Create the bind groups for both main textures
86///
87/// We will pick the correct one in the encoding system
88fn prepare_bind_groups(
89    mut commands: Commands,
90    mut views: Query<(Entity, &ViewTarget, Option<&mut PostProcessBindGroups>)>,
91    post_process_pipeline: Option<Res<PostProcessPipeline>>,
92    pipeline_cache: Res<PipelineCache>,
93    settings_uniforms: Res<ComponentUniforms<PostProcessSettings>>,
94    render_device: Res<RenderDevice>,
95) {
96    let Some(post_process_pipeline) = post_process_pipeline else {
97        return;
98    };
99    let Some(settings_binding) = settings_uniforms.uniforms().binding() else {
100        return;
101    };
102
103    let create_bind_group = |texture: &TextureView| {
104        (
105            texture.id(),
106            render_device.create_bind_group(
107                "post_process_bind_group",
108                &pipeline_cache.get_bind_group_layout(&post_process_pipeline.layout),
109                &BindGroupEntries::sequential((
110                    texture,
111                    &post_process_pipeline.sampler,
112                    settings_binding.clone(),
113                )),
114            ),
115        )
116    };
117
118    for (entity, view_target, mut maybe_bind_groups) in &mut views {
119        let main_texture_view = view_target.main_texture_view();
120        let main_texture_other_view = view_target.main_texture_other_view();
121
122        // Only update the cached bind groups if the main texture has changed
123        if let Some(bind_groups) = &mut maybe_bind_groups {
124            if bind_groups.a.0 != main_texture_view.id() {
125                bind_groups.a = create_bind_group(main_texture_view);
126            }
127            if bind_groups.b.0 != main_texture_other_view.id() {
128                bind_groups.b = create_bind_group(main_texture_other_view);
129            }
130        } else {
131            // Create the bind groups and add them to the view
132            commands.entity(entity).insert(PostProcessBindGroups {
133                a: create_bind_group(main_texture_view),
134                b: create_bind_group(main_texture_other_view),
135            });
136        }
137    }
138}
139
140fn post_process_system(
141    view: ViewQuery<(
142        &ViewTarget,
143        &DynamicUniformIndex<PostProcessSettings>,
144        &PostProcessBindGroups,
145    )>,
146    post_process_pipeline: Option<Res<PostProcessPipeline>>,
147    pipeline_cache: Res<PipelineCache>,
148    mut ctx: RenderContext,
149) {
150    let Some(post_process_pipeline) = post_process_pipeline else {
151        return;
152    };
153
154    let (view_target, settings_index, bind_groups) = view.into_inner();
155
156    let Some(pipeline) = pipeline_cache.get_render_pipeline(post_process_pipeline.pipeline_id)
157    else {
158        return;
159    };
160
161    // This will start a new "post process write", obtaining two texture
162    // views from the view target - a `source` and a `destination`.
163    // `source` is the "current" main texture and you _must_ write into
164    // `destination` because calling `post_process_write()` on the
165    // [`ViewTarget`] will internally flip the [`ViewTarget`]'s main
166    // texture to the `destination` texture. Failing to do so will cause
167    // the current main texture information to be lost.
168    let post_process = view_target.post_process_write();
169
170    // We prepared the bind group ahead of time but now we need to make sure
171    // we pick the bind group associated with the source texture
172    let (_, bind_group) = if bind_groups.a.0 == post_process.source.id() {
173        &bind_groups.a
174    } else {
175        &bind_groups.b
176    };
177
178    let mut render_pass = ctx
179        .command_encoder()
180        .begin_render_pass(&RenderPassDescriptor {
181            label: Some("post_process_pass"),
182            color_attachments: &[Some(RenderPassColorAttachment {
183                // We need to specify the post process destination view here
184                // to make sure we write to the appropriate texture.
185                view: post_process.destination,
186                depth_slice: None,
187                resolve_target: None,
188                ops: Operations::default(),
189            })],
190            depth_stencil_attachment: None,
191            timestamp_writes: None,
192            occlusion_query_set: None,
193            multiview_mask: None,
194        });
195
196    render_pass.set_pipeline(pipeline);
197    // By passing in the index of the post process settings on this view, we ensure
198    // that in the event that multiple settings were sent to the GPU (as would be the
199    // case with multiple cameras), we use the correct one.
200    render_pass.set_bind_group(0, bind_group, &[settings_index.index()]);
201    render_pass.draw(0..3, 0..1);
202}
203
204// This contains global data used by the render pipeline. This will be created once on startup.
205#[derive(Resource)]
206struct PostProcessPipeline {
207    layout: BindGroupLayoutDescriptor,
208    sampler: Sampler,
209    pipeline_id: CachedRenderPipelineId,
210}
211
212fn init_post_process_pipeline(
213    mut commands: Commands,
214    render_device: Res<RenderDevice>,
215    asset_server: Res<AssetServer>,
216    fullscreen_shader: Res<FullscreenShader>,
217    pipeline_cache: Res<PipelineCache>,
218) {
219    // We need to define the bind group layout used for our pipeline
220    let layout = BindGroupLayoutDescriptor::new(
221        "post_process_bind_group_layout",
222        &BindGroupLayoutEntries::sequential(
223            // The layout entries will only be visible in the fragment stage
224            ShaderStages::FRAGMENT,
225            (
226                // The screen texture
227                texture_2d(TextureSampleType::Float { filterable: true }),
228                // The sampler that will be used to sample the screen texture
229                sampler(SamplerBindingType::Filtering),
230                // The settings uniform that will control the effect
231                uniform_buffer::<PostProcessSettings>(true),
232            ),
233        ),
234    );
235    // We can create the sampler here since it won't change at runtime and doesn't depend on the view
236    let sampler = render_device.create_sampler(&SamplerDescriptor::default());
237
238    // Get the shader handle
239    let shader = asset_server.load(SHADER_ASSET_PATH);
240    // This will setup a fullscreen triangle for the vertex state.
241    let vertex_state = fullscreen_shader.to_vertex_state();
242    let pipeline_id = pipeline_cache
243        // This will add the pipeline to the cache and queue its creation
244        .queue_render_pipeline(RenderPipelineDescriptor {
245            label: Some("post_process_pipeline".into()),
246            layout: vec![layout.clone()],
247            vertex: vertex_state,
248            fragment: Some(FragmentState {
249                shader,
250                // Make sure this matches the entry point of your shader.
251                // It can be anything as long as it matches here and in the shader.
252                targets: vec![Some(ColorTargetState {
253                    format: TextureFormat::Rgba8UnormSrgb,
254                    blend: None,
255                    write_mask: ColorWrites::ALL,
256                })],
257                ..default()
258            }),
259            ..default()
260        });
261    commands.insert_resource(PostProcessPipeline {
262        layout,
263        sampler,
264        pipeline_id,
265    });
266}
267
268// This is the component that will get passed to the shader
269#[derive(Component, Default, Clone, Copy, ExtractComponent, ShaderType)]
270#[extract_app(RenderApp)]
271struct PostProcessSettings {
272    intensity: f32,
273    // WebGL2 structs must be 16 byte aligned.
274    #[cfg(feature = "webgl2")]
275    _webgl2_padding: Vec3,
276}
277
278/// Set up a simple 3D scene
279fn setup(
280    mut commands: Commands,
281    mut meshes: ResMut<Assets<Mesh>>,
282    mut materials: ResMut<Assets<StandardMaterial>>,
283) {
284    // camera
285    // Make sure you change the TextureFormat of the ColorTargetState
286    // if you enable Hdr directly or through features like Bloom.
287    commands.spawn((
288        Camera3d::default(),
289        Transform::from_translation(Vec3::new(0.0, 0.0, 5.0)).looking_at(Vec3::default(), Vec3::Y),
290        Camera {
291            clear_color: Color::WHITE.into(),
292            ..default()
293        },
294        // Add the setting to the camera.
295        // This component is also used to determine on which camera to run the post processing effect.
296        PostProcessSettings {
297            intensity: 0.02,
298            ..default()
299        },
300    ));
301
302    // cube
303    commands.spawn((
304        Mesh3d(meshes.add(Cuboid::default())),
305        MeshMaterial3d(materials.add(Color::srgb(0.8, 0.7, 0.6))),
306        Transform::from_xyz(0.0, 0.5, 0.0),
307        Rotates,
308    ));
309    // light
310    commands.spawn(DirectionalLight {
311        illuminance: 1_000.,
312        ..default()
313    });
314}
315
316#[derive(Component)]
317struct Rotates;
318
319/// Rotates any entity around the x and y axis
320fn rotate(time: Res<Time>, mut query: Query<&mut Transform, With<Rotates>>) {
321    for mut transform in &mut query {
322        transform.rotate_x(0.55 * time.delta_secs());
323        transform.rotate_z(0.15 * time.delta_secs());
324    }
325}
326
327// Change the intensity over time to show that the effect is controlled from the main world
328fn update_settings(mut settings: Query<&mut PostProcessSettings>, time: Res<Time>) {
329    for mut setting in &mut settings {
330        let mut intensity = ops::sin(time.elapsed_secs());
331        // Make it loop periodically
332        intensity = ops::sin(intensity);
333        // Remap it to 0..1 because the intensity can't be negative
334        intensity = intensity * 0.5 + 0.5;
335        // Scale it to a more reasonable level
336        intensity *= 0.015;
337
338        // Set the intensity.
339        // This will then be extracted to the render world and uploaded to the GPU automatically by the [`UniformComponentPlugin`]
340        setting.intensity = intensity;
341    }
342}