bevy_solari 0.20.0

Provides raytraced lighting for Bevy Engine
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
use super::{
    bind_group::BindGroupCacheState, instances::InstanceState, tlas_build, BlasManager,
    RaytracingSceneBindings,
};
use bevy_asset::load_embedded_asset;
use bevy_ecs::{
    resource::Resource,
    system::{Res, ResMut},
    world::{FromWorld, World},
};
use bevy_render::{
    diagnostic::RecordDiagnostics,
    render_resource::{
        binding_types::{storage_buffer_read_only_sized, storage_buffer_sized},
        AccelerationStructureFlags, AccelerationStructureUpdateMode, BindGroup, BindGroupEntries,
        BindGroupLayoutDescriptor, BindGroupLayoutEntries, Buffer, BufferDescriptor, BufferId,
        BufferUsages, CachedComputePipelineId, CommandEncoderDescriptor, ComputePassDescriptor,
        ComputePipelineDescriptor, CreateTlasDescriptor, PipelineCache, ShaderStages, Tlas,
        TlasInstance,
    },
    renderer::{RenderContext, RenderDevice},
};
use bevy_utils::{default, once};
use tracing::{info_span, warn};
use wgpu::{BufferTransition, BufferUses};

/// Compute pipeline that writes per-slot transforms and BLAS addresses into TLAS descriptors.
#[derive(Resource)]
pub struct TlasInstanceSetupPipeline {
    pub layout: BindGroupLayoutDescriptor,
    pub id: Option<CachedComputePipelineId>,
}

impl FromWorld for TlasInstanceSetupPipeline {
    fn from_world(world: &mut World) -> Self {
        let layout = BindGroupLayoutDescriptor::new(
            "tlas_instance_setup_bind_group_layout",
            &BindGroupLayoutEntries::sequential(
                ShaderStages::COMPUTE,
                (
                    storage_buffer_read_only_sized(false, None),
                    storage_buffer_read_only_sized(false, None),
                    storage_buffer_sized(false, None),
                ),
            ),
        );

        if tlas_build::resolve(world.resource::<RenderDevice>()).is_none() {
            return Self { layout, id: None };
        }

        let shader = load_embedded_asset!(world, "setup_tlas_instances.wesl");
        let id =
            world
                .resource::<PipelineCache>()
                .queue_compute_pipeline(ComputePipelineDescriptor {
                    label: Some("tlas_instance_setup_pipeline".into()),
                    layout: vec![layout.clone()],
                    shader,
                    entry_point: Some("setup_tlas_instances".into()),
                    ..default()
                });

        Self {
            layout,
            id: Some(id),
        }
    }
}

/// Smallest TLAS instance capacity handed out, and the floor [`tlas_capacity_for`] grows from.
const TLAS_MIN_CAPACITY: u32 = 128;

/// Width of a TLAS instance's custom data in both Vulkan (`instanceCustomIndex`) and DXR
/// (`InstanceID`). `setup_tlas_instances.wgsl` writes instance slots into that field, so they have
/// to fit.
const TLAS_CUSTOM_DATA_BITS: u32 = 24;

/// Instance capacity to allocate to hold `instance_count` slots.
///
/// Grows geometrically rather than by a fixed step, so reallocations over a scene's lifetime are
/// logarithmic in its size. Being a pure function of the count keeps the descriptor buffer and both
/// TLASes on the same capacity curve without having to coordinate their previous sizes.
fn tlas_capacity_for(instance_count: u32) -> u32 {
    let mut capacity = TLAS_MIN_CAPACITY;
    while capacity < instance_count {
        capacity = capacity.saturating_add(capacity.div_ceil(2));
    }
    capacity
}

/// The TLAS and its capacity state, double buffered only while something traces last frame's.
///
/// An acceleration structure is only bindable once it has been built. Nothing here owns a BLAS:
/// keeping the ones a built structure points at alive is [`BlasManager`]'s job, which defers every
/// retirement until [`BlasManager::note_tlas_build`] has seen both structures rebuilt.
pub struct TlasState {
    /// The backend to build through, or `None` to go through `wgpu-core`.
    raw: Option<&'static dyn tlas_build::RawTlasBackend>,
    /// Whether last frame's structure is being retained alongside this frame's.
    double_buffered: bool,
    /// Alternating current/previous acceleration structures. The second is `None` while single
    /// buffered.
    pub structures: [Option<Tlas>; 2],
    capacity: [u32; 2],
    /// Whether each structure has had a build recorded since its latest allocation.
    pub built: [bool; 2],
    /// Which of [`Self::structures`] this frame builds into. The other holds last frame's.
    pub current_index: usize,
    pub instance_descriptors: Option<Buffer>,
    instance_descriptor_capacity: u32,
    pub scratch: Option<Buffer>,
    scratch_capacity: u64,
    scratch_sized_for: u32,
    pub instance_setup_bind_group: Option<BindGroup>,
    instance_setup_buffer_ids: Option<[BufferId; 3]>,
}

impl TlasState {
    pub fn new(render_device: &RenderDevice) -> Self {
        Self {
            raw: tlas_build::resolve(render_device),
            double_buffered: false,
            structures: [None, None],
            capacity: [0, 0],
            built: [false, false],
            current_index: 0,
            instance_descriptors: None,
            instance_descriptor_capacity: 0,
            scratch: None,
            scratch_capacity: 0,
            scratch_sized_for: 0,
            instance_setup_bind_group: None,
            instance_setup_buffer_ids: None,
        }
    }

    /// Whether builds go through [`tlas_build`]'s raw path rather than `wgpu-core`.
    pub fn uses_raw_build(&self) -> bool {
        self.raw.is_some()
    }

    /// Brings the acceleration structure up to date with this frame's changes.
    ///
    /// While `double_buffered`, the two structures alternate: this frame's is rebuilt, and last
    /// frame's stays intact so the shaders can trace against it. Otherwise the single structure is
    /// rebuilt in place, which is safe because the build is recorded before any pass that traces
    /// it and nothing reads last frame's contents.
    ///
    /// `build_ready` reports whether this frame will be able to record a build. When false nothing
    /// is allocated or swapped, though an unused structure is still released.
    pub fn advance(
        &mut self,
        instances: &InstanceState,
        bind_groups: &mut BindGroupCacheState,
        render_device: &RenderDevice,
        build_ready: bool,
        double_buffered: bool,
    ) {
        let _span = info_span!("advance_tlas").entered();

        // Free the retained structure as soon as nothing wants it
        self.set_double_buffered(double_buffered, bind_groups);

        // An empty scene must not allocate an unbuilt TLAS that could resurface as a later
        // previous-frame entry
        if !build_ready || instances.slots.high_water_mark() == 0 {
            return;
        }

        debug_assert!(
            instances.slots.high_water_mark() < 1 << TLAS_CUSTOM_DATA_BITS,
            "instance slot count {} does not fit in a TLAS instance's custom data",
            instances.slots.high_water_mark()
        );

        // Secure every build input before committing the swap. Neither buffer exists on the
        // wgpu-core path, which builds from instances written into the TLAS itself.
        let instance_count = instances.slots.high_water_mark();
        if self.raw.is_some() {
            self.reserve_instance_descriptors(instance_count, render_device);
            self.reserve_tlas_scratch(render_device);
            if self.instance_descriptors.is_none() || self.scratch.is_none() {
                return;
            }
        }

        if self.double_buffered {
            self.current_index ^= 1;
        }
        let current_index = self.current_index;
        self.reserve_tlas(current_index, instance_count, render_device, bind_groups);
    }

    /// Switches between retaining last frame's structure and rebuilding a single one in place,
    /// dropping the structure that is no longer traced.
    fn set_double_buffered(
        &mut self,
        double_buffered: bool,
        bind_groups: &mut BindGroupCacheState,
    ) {
        if self.double_buffered != double_buffered {
            self.double_buffered = double_buffered;
            bind_groups.invalid = true;
        }
        if double_buffered {
            return;
        }

        // Keep building into whichever structure is already current, so the frame this turns off
        // does not have to allocate
        let previous_index = self.current_index ^ 1;
        if self.structures[previous_index].take().is_some() {
            self.capacity[previous_index] = 0;
            self.built[previous_index] = false;
            bind_groups.invalid = true;
        }
    }

    /// Whether the previous-frame slot's binding will still be correct next frame, and so whether
    /// a bind group using it can be cached.
    pub fn previous_binding_is_stable(&self) -> bool {
        !self.double_buffered || self.built[self.current_index ^ 1]
    }

    /// Rebuilds the setup shader's bind group only when one of its three buffers moves.
    pub fn update_instance_setup_bind_group(
        &mut self,
        instances: &InstanceState,
        render_device: &RenderDevice,
        pipeline_cache: &PipelineCache,
        pipeline: &TlasInstanceSetupPipeline,
    ) {
        let (Some(transforms), Some(blas_refs), Some(instances)) = (
            instances.transforms.buffer(),
            instances.blas_refs.buffer(),
            self.instance_descriptors.as_ref(),
        ) else {
            self.instance_setup_bind_group = None;
            return;
        };

        let ids = [transforms.id(), blas_refs.id(), instances.id()];
        if self.instance_setup_bind_group.is_some() && self.instance_setup_buffer_ids == Some(ids) {
            return;
        }

        let layout = pipeline_cache.get_bind_group_layout(&pipeline.layout);
        self.instance_setup_bind_group = Some(render_device.create_bind_group(
            "tlas_instance_setup_bind_group",
            &layout,
            &BindGroupEntries::sequential((
                transforms.as_entire_binding(),
                blas_refs.as_entire_binding(),
                instances.as_entire_binding(),
            )),
        ));
        self.instance_setup_buffer_ids = Some(ids);
    }

    /// Makes sure the instance descriptor buffer covers every stable slot.
    fn reserve_instance_descriptors(&mut self, needed: u32, render_device: &RenderDevice) {
        if self.instance_descriptors.is_some() && needed <= self.instance_descriptor_capacity {
            return;
        }

        let capacity = tlas_capacity_for(needed);
        self.instance_descriptors = Some(render_device.create_buffer(&BufferDescriptor {
            label: Some("solari_tlas_instance_descriptors"),
            size: u64::from(capacity) * tlas_build::INSTANCE_DESCRIPTOR_SIZE,
            usage: BufferUsages::STORAGE | BufferUsages::TLAS_INPUT,
            mapped_at_creation: false,
        }));
        self.instance_descriptor_capacity = capacity;
        self.instance_setup_bind_group = None;
    }

    /// Makes sure the build scratch buffer is big enough for this frame's descriptor capacity.
    fn reserve_tlas_scratch(&mut self, render_device: &RenderDevice) {
        let capacity = self.instance_descriptor_capacity;
        if self.scratch.is_some() && capacity <= self.scratch_sized_for {
            return;
        }

        let (Some(backend), Some(instances)) = (self.raw, self.instance_descriptors.as_ref())
        else {
            return;
        };
        let Some(needed) = backend.scratch_size(render_device, instances, capacity) else {
            return;
        };
        self.scratch_sized_for = capacity;

        if self.scratch.is_some() && needed <= self.scratch_capacity {
            return;
        }

        // The setup transition lets wgpu retain an outgrown scratch buffer until in-flight work
        // releases it
        self.scratch = backend.create_scratch_buffer(render_device, needed);
        self.scratch_capacity = if self.scratch.is_some() { needed } else { 0 };
    }

    /// Makes sure one of the two structures can hold every instance slot, leaving the other alone.
    fn reserve_tlas(
        &mut self,
        current_index: usize,
        needed: u32,
        render_device: &RenderDevice,
        bind_groups: &mut BindGroupCacheState,
    ) {
        if self.structures[current_index].is_some() && needed <= self.capacity[current_index] {
            return;
        }

        let capacity = tlas_capacity_for(needed);
        self.structures[current_index] = Some(render_device.wgpu_device().create_tlas(
            &CreateTlasDescriptor {
                label: Some("tlas"),
                flags: AccelerationStructureFlags::PREFER_FAST_TRACE,
                update_mode: AccelerationStructureUpdateMode::Build,
                max_instances: capacity,
            },
        ));
        self.capacity[current_index] = capacity;
        self.built[current_index] = false;
        bind_groups.invalid = true;
    }
}

/// Records this frame's TLAS build into the render graph's command encoder.
pub fn build_raytracing_tlas(
    mut bindings: ResMut<RaytracingSceneBindings>,
    mut blas_manager: ResMut<BlasManager>,
    pipeline_cache: Res<PipelineCache>,
    pipeline: Res<TlasInstanceSetupPipeline>,
    mut render_context: RenderContext,
) {
    let bindings = &mut *bindings;
    let current_index = bindings.tlas.current_index;

    let built = match bindings.tlas.raw {
        Some(backend) => {
            setup_tlas_instances(bindings, &pipeline_cache, &pipeline, &mut render_context)
                && build_tlas_raw(bindings, backend, &mut render_context)
        }
        None => build_tlas_through_wgpu_core(bindings, &blas_manager, &mut render_context),
    };

    if built {
        bindings.tlas.built[current_index] = true;
        // This structure no longer points at whatever was retired before it, which is what lets the
        // oldest retirements go
        blas_manager.note_tlas_build();
    }
}

/// Writes this frame's TLAS instance descriptors on the GPU, ready for [`build_tlas_raw`].
///
/// Returns whether the pass was recorded. The build has nothing to read when it wasn't.
fn setup_tlas_instances(
    bindings: &mut RaytracingSceneBindings,
    pipeline_cache: &PipelineCache,
    pipeline: &TlasInstanceSetupPipeline,
    render_context: &mut RenderContext,
) -> bool {
    if bindings.tlas.structures[bindings.tlas.current_index].is_none() {
        return false;
    }

    let (Some(bind_group), Some(compute_pipeline), Some(instances), Some(scratch)) = (
        bindings.tlas.instance_setup_bind_group.as_ref(),
        pipeline
            .id
            .and_then(|id| pipeline_cache.get_compute_pipeline(id)),
        bindings.tlas.instance_descriptors.as_ref(),
        bindings.tlas.scratch.as_ref(),
    ) else {
        // `advance` secures all of these before it allocates the TLAS this is setting up for
        once!(warn!(
            "TLAS allocated but its instance setup pass could not be recorded: bind group={}, \
             descriptors={}, scratch={}",
            bindings.tlas.instance_setup_bind_group.is_some(),
            bindings.tlas.instance_descriptors.is_some(),
            bindings.tlas.scratch.is_some(),
        ));
        return false;
    };

    let slot_count = bindings.instances.slots.high_water_mark();
    if slot_count == 0 {
        return false;
    }

    let diagnostics = render_context.diagnostic_recorder();
    let diagnostics = diagnostics.as_deref();
    let command_encoder = render_context.command_encoder();
    let time_span = diagnostics.time_span(command_encoder, "setup_tlas_instances");
    {
        let mut pass = command_encoder.begin_compute_pass(&ComputePassDescriptor {
            label: Some("setup_tlas_instances"),
            timestamp_writes: None,
        });
        pass.set_pipeline(compute_pipeline);
        pass.set_bind_group(0, bind_group, &[]);
        pass.dispatch_workgroups(slot_count.div_ceil(64), 1, 1);
    }
    time_span.end(command_encoder);

    command_encoder.transition_resources(
        [
            BufferTransition {
                buffer: &**instances,
                state: BufferUses::TOP_LEVEL_ACCELERATION_STRUCTURE_INPUT,
            },
            BufferTransition {
                buffer: &**scratch,
                state: BufferUses::ACCELERATION_STRUCTURE_SCRATCH,
            },
        ]
        .into_iter(),
        core::iter::empty(),
    );

    true
}

/// Builds from the descriptors [`setup_tlas_instances`] wrote, straight through `wgpu_hal`.
fn build_tlas_raw(
    bindings: &mut RaytracingSceneBindings,
    backend: &dyn tlas_build::RawTlasBackend,
    render_context: &mut RenderContext,
) -> bool {
    let current_index = bindings.tlas.current_index;
    let (Some(tlas), Some(instances), Some(scratch)) = (
        bindings.tlas.structures[current_index].as_mut(),
        bindings.tlas.instance_descriptors.as_ref(),
        bindings.tlas.scratch.as_ref(),
    ) else {
        once!(warn!(
            "TLAS instances were set up but not built: structure={}, descriptors={}, scratch={}",
            bindings.tlas.structures[current_index].is_some(),
            bindings.tlas.instance_descriptors.is_some(),
            bindings.tlas.scratch.is_some(),
        ));
        return false;
    };

    let render_device = render_context.render_device().clone();
    let diagnostics = render_context.diagnostic_recorder();
    let diagnostics = diagnostics.as_deref();
    let time_span = diagnostics.time_span(render_context.command_encoder(), "tlas_build");

    let mut command_encoder = render_device.create_command_encoder(&CommandEncoderDescriptor {
        label: Some("tlas_build_command_encoder"),
    });
    let built = backend.build_tlas(
        &mut command_encoder,
        tlas,
        instances,
        bindings.instances.slots.high_water_mark(),
        scratch,
    );
    render_context.add_command_buffer(command_encoder.finish());

    time_span.end(render_context.command_encoder());
    if !built {
        once!(warn!(
            "TLAS build recorded nothing; the resolved backend does not own the build resources."
        ));
    }
    built
}

/// Builds from instance descriptors filled in on the CPU, through `wgpu-core`.
///
/// The portable path, and the only one Metal can use. It costs work per instance every frame rather
/// than only for the ones that moved, and `wgpu-core` then re-derives the whole build from what it
/// is handed here, which is the cost [`tlas_build`] exists to avoid.
fn build_tlas_through_wgpu_core(
    bindings: &mut RaytracingSceneBindings,
    blas_manager: &BlasManager,
    render_context: &mut RenderContext,
) -> bool {
    // An empty scene leaves the index unswapped, so whatever is here belongs to an earlier frame
    // and is still being traced as the previous one
    if bindings.instances.slots.high_water_mark() == 0 {
        return false;
    }
    let current_index = bindings.tlas.current_index;
    let Some(tlas) = bindings.tlas.structures[current_index].as_mut() else {
        return false;
    };

    {
        let _span = info_span!("fill_tlas_instances").entered();

        // The TLAS outlives the frame, so a slot freed or deactivated since its last build still
        // holds that build's instance and has to be cleared before this frame's are written
        let capacity = tlas.get().len();
        tlas[0..capacity].iter_mut().for_each(|entry| *entry = None);

        for (slot, mesh, transform) in bindings.instances.drawable() {
            // A mesh can lose its acceleration structure after the instance resolved against it,
            // which leaves the slot with nothing to point at for a frame
            let Some(blas) = blas_manager.get(&mesh) else {
                continue;
            };
            tlas[slot as usize] = Some(TlasInstance::new(blas, transform, slot, 0xFF));
        }
    }

    let diagnostics = render_context.diagnostic_recorder();
    let diagnostics = diagnostics.as_deref();
    let command_encoder = render_context.command_encoder();
    let time_span = diagnostics.time_span(command_encoder, "tlas_build");
    command_encoder.build_acceleration_structures(&[], [&*tlas]);
    time_span.end(command_encoder);

    true
}

#[cfg(test)]
mod tests {
    use super::{tlas_capacity_for, TLAS_MIN_CAPACITY};

    #[test]
    fn tlas_capacity_grows_geometrically_at_boundaries() {
        assert_eq!(tlas_capacity_for(0), TLAS_MIN_CAPACITY);
        assert_eq!(tlas_capacity_for(TLAS_MIN_CAPACITY), TLAS_MIN_CAPACITY);
        assert_eq!(tlas_capacity_for(TLAS_MIN_CAPACITY + 1), 192);
        assert_eq!(tlas_capacity_for(192), 192);
        assert_eq!(tlas_capacity_for(193), 288);
    }
}