use wgpu::TextureFormat;
pub use crate::render::wgsl::OIT_WGSL;
pub const OIT_MAX_LAYERS: u32 = 16;
pub const OIT_DEFAULT_LAYERS: u32 = 8;
pub const OIT_NODE_SIZE: u64 = 12;
#[repr(C)]
#[derive(Clone, Copy, Debug, Default, PartialEq, bytemuck::Pod, bytemuck::Zeroable)]
pub struct OitVertex {
pub pos: [f32; 2],
pub depth: f32,
pub _pad: f32,
pub color: [f32; 4],
}
pub const OIT_VERTEX_STRIDE: u64 = 32;
#[repr(C)]
#[derive(Clone, Copy, Debug, Default, bytemuck::Pod, bytemuck::Zeroable)]
struct OitUniform {
viewport: [f32; 2],
dims: [u32; 2],
background: [f32; 4],
capacity: u32,
_pad: [u32; 3],
}
#[derive(Clone, Debug, Default)]
pub struct OitBatch {
verts: Vec<OitVertex>,
}
impl OitBatch {
#[must_use]
pub fn new() -> Self {
Self { verts: Vec::new() }
}
pub fn clear(&mut self) {
self.verts.clear();
}
pub fn push_tri(
&mut self,
a: egui::Pos2,
b: egui::Pos2,
c: egui::Pos2,
color: [f32; 4],
depth: f32,
) {
for p in [a, b, c] {
self.verts.push(OitVertex { pos: [p.x, p.y], depth, _pad: 0.0, color });
}
}
pub fn push_quad(&mut self, rect: egui::Rect, color: [f32; 4], depth: f32) {
let (a, b) = (rect.min, rect.max);
let tr = egui::pos2(b.x, a.y);
let bl = egui::pos2(a.x, b.y);
self.push_tri(a, tr, b, color, depth);
self.push_tri(a, b, bl, color, depth);
}
#[must_use]
pub fn vertices(&self) -> &[OitVertex] {
&self.verts
}
#[must_use]
pub fn len(&self) -> usize {
self.verts.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.verts.is_empty()
}
#[must_use]
pub fn permuted(&self, order: &[usize]) -> Self {
let mut out = Self { verts: Vec::with_capacity(self.verts.len()) };
for &t in order {
let base = t * 3;
if base + 2 < self.verts.len() {
out.verts.extend_from_slice(&self.verts[base..base + 3]);
}
}
out
}
#[must_use]
pub fn tri_count(&self) -> usize {
self.verts.len() / 3
}
}
pub struct OitPass {
gather: wgpu::RenderPipeline,
resolve: wgpu::RenderPipeline,
bgl: wgpu::BindGroupLayout,
uniform: wgpu::Buffer,
heads: Option<wgpu::Buffer>,
nodes: Option<wgpu::Buffer>,
alloc: wgpu::Buffer,
bind: Option<wgpu::BindGroup>,
verts: Option<wgpu::Buffer>,
vert_cap: usize,
vertex_count: u32,
size: (u32, u32),
capacity: u32,
layers: u32,
}
impl OitPass {
#[must_use]
pub fn supported(adapter: &wgpu::Adapter) -> bool {
adapter
.get_downlevel_capabilities()
.flags
.contains(wgpu::DownlevelFlags::FRAGMENT_WRITABLE_STORAGE)
}
#[must_use]
pub fn new(device: &wgpu::Device, target_format: TextureFormat) -> Self {
let module = device.create_shader_module(wgpu::ShaderModuleDescriptor {
label: Some("l0_oit"),
source: wgpu::ShaderSource::Wgsl(super::wgsl(OIT_WGSL).into()),
});
let storage = |binding: u32| wgpu::BindGroupLayoutEntry {
binding,
visibility: wgpu::ShaderStages::FRAGMENT,
ty: wgpu::BindingType::Buffer {
ty: wgpu::BufferBindingType::Storage { read_only: false },
has_dynamic_offset: false,
min_binding_size: None,
},
count: None,
};
let bgl = device.create_bind_group_layout(&wgpu::BindGroupLayoutDescriptor {
label: Some("l0_oit_bgl"),
entries: &[
wgpu::BindGroupLayoutEntry {
binding: 0,
visibility: wgpu::ShaderStages::VERTEX_FRAGMENT,
ty: wgpu::BindingType::Buffer {
ty: wgpu::BufferBindingType::Uniform,
has_dynamic_offset: false,
min_binding_size: None,
},
count: None,
},
storage(1),
storage(2),
storage(3),
],
});
let pll = device.create_pipeline_layout(&wgpu::PipelineLayoutDescriptor {
label: Some("l0_oit_pll"),
bind_group_layouts: &[Some(&bgl)],
immediate_size: 0,
});
let gather = device.create_render_pipeline(&wgpu::RenderPipelineDescriptor {
label: Some("l0_oit_gather"),
layout: Some(&pll),
vertex: wgpu::VertexState {
module: &module,
entry_point: Some("oit_gather_vs"),
compilation_options: Default::default(),
buffers: &[wgpu::VertexBufferLayout {
array_stride: OIT_VERTEX_STRIDE,
step_mode: wgpu::VertexStepMode::Vertex,
attributes: &[
wgpu::VertexAttribute {
format: wgpu::VertexFormat::Float32x2,
offset: 0,
shader_location: 0,
},
wgpu::VertexAttribute {
format: wgpu::VertexFormat::Float32,
offset: 8,
shader_location: 1,
},
wgpu::VertexAttribute {
format: wgpu::VertexFormat::Float32x4,
offset: 16,
shader_location: 2,
},
],
}],
},
primitive: wgpu::PrimitiveState {
topology: wgpu::PrimitiveTopology::TriangleList,
cull_mode: None,
..Default::default()
},
depth_stencil: None,
multisample: wgpu::MultisampleState::default(),
fragment: Some(wgpu::FragmentState {
module: &module,
entry_point: Some("oit_gather_fs"),
compilation_options: Default::default(),
targets: &[Some(wgpu::ColorTargetState {
format: target_format,
blend: None,
write_mask: wgpu::ColorWrites::empty(),
})],
}),
multiview_mask: None,
cache: None,
});
let resolve = device.create_render_pipeline(&wgpu::RenderPipelineDescriptor {
label: Some("l0_oit_resolve"),
layout: Some(&pll),
vertex: wgpu::VertexState {
module: &module,
entry_point: Some("oit_resolve_vs"),
compilation_options: Default::default(),
buffers: &[],
},
primitive: wgpu::PrimitiveState {
topology: wgpu::PrimitiveTopology::TriangleList,
..Default::default()
},
depth_stencil: None,
multisample: wgpu::MultisampleState::default(),
fragment: Some(wgpu::FragmentState {
module: &module,
entry_point: Some("oit_resolve_fs"),
compilation_options: Default::default(),
targets: &[Some(wgpu::ColorTargetState {
format: target_format,
blend: None,
write_mask: wgpu::ColorWrites::ALL,
})],
}),
multiview_mask: None,
cache: None,
});
Self {
gather,
resolve,
bgl,
uniform: device.create_buffer(&wgpu::BufferDescriptor {
label: Some("l0_oit_uniform"),
size: std::mem::size_of::<OitUniform>() as u64,
usage: wgpu::BufferUsages::UNIFORM | wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
}),
heads: None,
nodes: None,
alloc: device.create_buffer(&wgpu::BufferDescriptor {
label: Some("l0_oit_alloc"),
size: 4,
usage: wgpu::BufferUsages::STORAGE
| wgpu::BufferUsages::COPY_DST
| wgpu::BufferUsages::COPY_SRC,
mapped_at_creation: false,
}),
bind: None,
verts: None,
vert_cap: 0,
vertex_count: 0,
size: (0, 0),
capacity: 0,
layers: 0,
}
}
pub fn ensure(&mut self, device: &wgpu::Device, w: u32, h: u32, layers: u32) {
let w = w.max(1);
let h = h.max(1);
let layers = layers.max(1);
if self.size == (w, h) && self.layers == layers && self.heads.is_some() {
return;
}
let pixels = u64::from(w) * u64::from(h);
let heads = device.create_buffer(&wgpu::BufferDescriptor {
label: Some("l0_oit_heads"),
size: pixels * 4,
usage: wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
});
let capacity = pixels * u64::from(layers);
let nodes = device.create_buffer(&wgpu::BufferDescriptor {
label: Some("l0_oit_nodes"),
size: capacity * OIT_NODE_SIZE,
usage: wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
});
self.bind = Some(device.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some("l0_oit_bind"),
layout: &self.bgl,
entries: &[
wgpu::BindGroupEntry { binding: 0, resource: self.uniform.as_entire_binding() },
wgpu::BindGroupEntry { binding: 1, resource: heads.as_entire_binding() },
wgpu::BindGroupEntry { binding: 2, resource: nodes.as_entire_binding() },
wgpu::BindGroupEntry { binding: 3, resource: self.alloc.as_entire_binding() },
],
}));
self.heads = Some(heads);
self.nodes = Some(nodes);
self.size = (w, h);
self.layers = layers;
self.capacity = u32::try_from(capacity).unwrap_or(u32::MAX);
}
pub fn set_frame(&self, queue: &wgpu::Queue, background: [f32; 4]) {
let (w, h) = self.size;
queue.write_buffer(
&self.uniform,
0,
bytemuck::bytes_of(&OitUniform {
viewport: [w as f32, h as f32],
dims: [w, h],
background,
capacity: self.capacity,
_pad: [0; 3],
}),
);
}
pub fn upload(&mut self, device: &wgpu::Device, queue: &wgpu::Queue, batch: &OitBatch) {
let verts = batch.vertices();
self.vertex_count = u32::try_from(verts.len()).unwrap_or(u32::MAX);
if verts.is_empty() {
return;
}
if self.vert_cap < verts.len() || self.verts.is_none() {
let cap = verts.len().next_power_of_two();
self.verts = Some(device.create_buffer(&wgpu::BufferDescriptor {
label: Some("l0_oit_verts"),
size: cap as u64 * OIT_VERTEX_STRIDE,
usage: wgpu::BufferUsages::VERTEX | wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
}));
self.vert_cap = cap;
}
if let Some(vb) = &self.verts {
queue.write_buffer(vb, 0, bytemuck::cast_slice(verts));
}
}
pub fn reset(&self, encoder: &mut wgpu::CommandEncoder) {
if let Some(heads) = &self.heads {
encoder.clear_buffer(heads, 0, None);
}
encoder.clear_buffer(&self.alloc, 0, None);
}
pub fn record(&self, encoder: &mut wgpu::CommandEncoder, target: &wgpu::TextureView) -> u32 {
let Some(bind) = &self.bind else { return 0 };
self.reset(encoder);
if self.vertex_count > 0 {
if let Some(vb) = &self.verts {
let mut rp = encoder.begin_render_pass(&wgpu::RenderPassDescriptor {
label: Some("l0_oit_gather_pass"),
color_attachments: &[Some(wgpu::RenderPassColorAttachment {
view: target,
resolve_target: None,
depth_slice: None,
ops: wgpu::Operations {
load: wgpu::LoadOp::Clear(wgpu::Color::TRANSPARENT),
store: wgpu::StoreOp::Store,
},
})],
depth_stencil_attachment: None,
timestamp_writes: None,
occlusion_query_set: None,
multiview_mask: None,
});
rp.set_pipeline(&self.gather);
rp.set_bind_group(0, bind, &[]);
rp.set_vertex_buffer(0, vb.slice(..));
rp.draw(0..self.vertex_count, 0..1);
}
}
{
let mut rp = encoder.begin_render_pass(&wgpu::RenderPassDescriptor {
label: Some("l0_oit_resolve_pass"),
color_attachments: &[Some(wgpu::RenderPassColorAttachment {
view: target,
resolve_target: None,
depth_slice: None,
ops: wgpu::Operations {
load: wgpu::LoadOp::Clear(wgpu::Color::TRANSPARENT),
store: wgpu::StoreOp::Store,
},
})],
depth_stencil_attachment: None,
timestamp_writes: None,
occlusion_query_set: None,
multiview_mask: None,
});
rp.set_pipeline(&self.resolve);
rp.set_bind_group(0, bind, &[]);
rp.draw(0..3, 0..1);
}
self.vertex_count
}
#[must_use]
pub fn fragment_count(&self, device: &wgpu::Device, queue: &wgpu::Queue) -> u32 {
let bytes = super::readback::read_buffer_range(device, queue, &self.alloc, 0, 4);
if bytes.len() < 4 {
return 0;
}
u32::from_le_bytes([bytes[0], bytes[1], bytes[2], bytes[3]])
}
#[must_use]
pub fn capacity(&self) -> u32 {
self.capacity
}
#[must_use]
pub fn size(&self) -> (u32, u32) {
self.size
}
#[must_use]
pub fn ready(&self) -> bool {
self.bind.is_some()
}
#[must_use]
pub fn vertex_count(&self) -> u32 {
self.vertex_count
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn vertex_layout_matches_the_shader() {
assert_eq!(std::mem::size_of::<OitVertex>() as u64, OIT_VERTEX_STRIDE);
assert_eq!(std::mem::offset_of!(OitVertex, pos), 0);
assert_eq!(std::mem::offset_of!(OitVertex, depth), 8);
assert_eq!(std::mem::offset_of!(OitVertex, color), 16);
assert_eq!(std::mem::size_of::<OitUniform>(), 48, "the uniform is 3 × 16 B");
}
#[test]
fn the_layer_count_constant_is_the_same_number_on_both_sides() {
let want = format!("const OIT_LAYERS: u32 = {OIT_MAX_LAYERS}u;");
assert!(
OIT_WGSL.contains(&want),
"OIT_MAX_LAYERS ({OIT_MAX_LAYERS}) must match the shader's OIT_LAYERS; looked for `{want}`"
);
}
#[test]
fn the_gather_neither_depth_tests_nor_culls() {
assert!(
OIT_WGSL.contains("o.clip = vec4<f32>(px_to_ndc(v.pos_px, u.viewport), 0.0, 1.0);"),
"the gather must emit a constant z — a varying z with a depth test would cull"
);
}
#[test]
fn a_quad_is_two_triangles_of_one_depth_over_its_rect() {
let mut b = OitBatch::new();
let r = egui::Rect::from_min_max(egui::pos2(10.0, 20.0), egui::pos2(30.0, 50.0));
b.push_quad(r, [0.25, 0.5, 0.75, 0.5], 3.5);
assert_eq!(b.len(), 6);
assert_eq!(b.tri_count(), 2);
assert!(b.vertices().iter().all(|v| v.depth == 3.5), "one depth");
assert!(b.vertices().iter().all(|v| v.color == [0.25, 0.5, 0.75, 0.5]), "one colour");
let xs: Vec<f32> = b.vertices().iter().map(|v| v.pos[0]).collect();
let ys: Vec<f32> = b.vertices().iter().map(|v| v.pos[1]).collect();
assert_eq!(xs.iter().cloned().fold(f32::MAX, f32::min), 10.0);
assert_eq!(xs.iter().cloned().fold(f32::MIN, f32::max), 30.0);
assert_eq!(ys.iter().cloned().fold(f32::MAX, f32::min), 20.0);
assert_eq!(ys.iter().cloned().fold(f32::MIN, f32::max), 50.0);
assert_eq!(b.vertices()[0].pos, [10.0, 20.0]);
assert_eq!(b.vertices()[2].pos, [30.0, 50.0]);
assert_eq!(b.vertices()[3].pos, [10.0, 20.0]);
assert_eq!(b.vertices()[4].pos, [30.0, 50.0]);
}
#[test]
fn permuting_reorders_whole_triangles_and_preserves_the_multiset() {
let mut b = OitBatch::new();
for i in 0..4u32 {
let x = i as f32 * 10.0;
b.push_quad(
egui::Rect::from_min_max(egui::pos2(x, 0.0), egui::pos2(x + 5.0, 5.0)),
[i as f32 / 4.0, 0.0, 0.0, 0.5],
i as f32,
);
}
let n = b.tri_count();
assert_eq!(n, 8);
let rev: Vec<usize> = (0..n).rev().collect();
let p = b.permuted(&rev);
assert_eq!(p.len(), b.len(), "same vertex count");
let groups = |x: &OitBatch| -> Vec<[OitVertex; 3]> {
x.vertices().chunks(3).map(|c| [c[0], c[1], c[2]]).collect()
};
let (og, pg) = (groups(&b), groups(&p));
assert_eq!(pg.len(), og.len());
for g in &pg {
assert!(og.contains(g), "each permuted triangle is an original triangle");
}
assert_ne!(pg, og, "the reversal must actually change the submission order");
assert_eq!(pg.iter().rev().cloned().collect::<Vec<_>>(), og, "…by reversing it");
}
#[test]
fn permuting_is_total() {
let mut b = OitBatch::new();
b.push_quad(egui::Rect::from_min_max(egui::pos2(0.0, 0.0), egui::pos2(4.0, 4.0)), [1.0; 4], 0.0);
assert_eq!(b.permuted(&[99]).len(), 0, "out-of-range triangle is dropped");
assert_eq!(b.permuted(&[]).len(), 0);
assert_eq!(b.permuted(&[0, 1]).len(), 6);
}
}