nightshade-renderer 0.57.0

GPU-driven wgpu renderer with a built-in frame graph.
use crate::config::SkinnedAnimationSnapshot;
use crate::skinning::SkinningCache;
use nightshade_ecs::Entity;
use wgpu::util::DeviceExt;

#[repr(C)]
#[derive(Copy, Clone, bytemuck::Pod, bytemuck::Zeroable)]
struct GpuSpringChain {
    joint_start: u32,
    joint_count: u32,
    descendant_start: u32,
    descendant_count: u32,
    state_start: u32,
    reset: u32,
    enabled: u32,
    pad0: u32,
    gravity: [f32; 4],
    stiffness: f32,
    damping: f32,
    dt: f32,
    pad1: f32,
}

struct SpringChainLayout {
    player: Entity,
    chain_index: usize,
    joint_start: u32,
    joint_count: u32,
    descendant_start: u32,
    descendant_count: u32,
    state_start: u32,
}

pub(super) struct SkinnedSpringGpu {
    pipeline: wgpu::ComputePipeline,
    bind_group_layout: wgpu::BindGroupLayout,
    bind_group: Option<wgpu::BindGroup>,
    chains_buffer: wgpu::Buffer,
    joint_indices_buffer: wgpu::Buffer,
    tip_descendants_buffer: wgpu::Buffer,
    rest_axes_buffer: wgpu::Buffer,
    bone_lengths_buffer: wgpu::Buffer,
    state_buffer: wgpu::Buffer,
    layouts: Vec<SpringChainLayout>,
    reset_pending: bool,
    cached_static_signature: u64,
    bone_transforms_generation: u64,
}

fn sized_buffer(device: &wgpu::Device, label: &str, size: u64) -> wgpu::Buffer {
    device.create_buffer(&wgpu::BufferDescriptor {
        label: Some(label),
        size: size.max(16),
        usage: wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_DST,
        mapped_at_creation: false,
    })
}

fn init_buffer(device: &wgpu::Device, label: &str, data: &[u8]) -> wgpu::Buffer {
    let fallback = [0u8; 16];
    device.create_buffer_init(&wgpu::util::BufferInitDescriptor {
        label: Some(label),
        contents: if data.is_empty() { &fallback } else { data },
        usage: wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_DST,
    })
}

impl SkinnedSpringGpu {
    pub(super) fn new(device: &wgpu::Device) -> Self {
        let shader = crate::wgpu::shader_compose::compile_wgsl(
            device,
            "spring_compute.wgsl",
            include_str!("../../../shaders/spring_compute.wgsl"),
        );

        let entries: Vec<wgpu::BindGroupLayoutEntry> = (0..7)
            .map(|binding| wgpu::BindGroupLayoutEntry {
                binding,
                visibility: wgpu::ShaderStages::COMPUTE,
                ty: wgpu::BindingType::Buffer {
                    ty: wgpu::BufferBindingType::Storage {
                        read_only: binding != 0 && binding != 6,
                    },
                    has_dynamic_offset: false,
                    min_binding_size: None,
                },
                count: None,
            })
            .collect();
        let bind_group_layout = device.create_bind_group_layout(&wgpu::BindGroupLayoutDescriptor {
            label: Some("Skinned Spring Bind Group Layout"),
            entries: &entries,
        });

        let pipeline_layout = device.create_pipeline_layout(&wgpu::PipelineLayoutDescriptor {
            label: Some("Skinned Spring Pipeline Layout"),
            bind_group_layouts: &[Some(&bind_group_layout)],
            immediate_size: 0,
        });

        let pipeline = device.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
            label: Some("Skinned Spring Pipeline"),
            layout: Some(&pipeline_layout),
            module: &shader,
            entry_point: Some("main"),
            compilation_options: Default::default(),
            cache: None,
        });

        Self {
            pipeline,
            bind_group_layout,
            bind_group: None,
            chains_buffer: sized_buffer(device, "Spring Chains", 16),
            joint_indices_buffer: sized_buffer(device, "Spring Joint Indices", 16),
            tip_descendants_buffer: sized_buffer(device, "Spring Tip Descendants", 16),
            rest_axes_buffer: sized_buffer(device, "Spring Rest Axes", 16),
            bone_lengths_buffer: sized_buffer(device, "Spring Bone Lengths", 16),
            state_buffer: sized_buffer(device, "Spring State", 32),
            layouts: Vec::new(),
            reset_pending: true,
            cached_static_signature: u64::MAX,
            bone_transforms_generation: u64::MAX,
        }
    }

    pub(super) fn chain_count(&self) -> u32 {
        self.layouts.len() as u32
    }

    pub(super) fn dispatch(&self, compute_pass: &mut wgpu::ComputePass) {
        if self.layouts.is_empty() {
            return;
        }
        if let Some(bind_group) = self.bind_group.as_ref() {
            compute_pass.set_pipeline(&self.pipeline);
            compute_pass.set_bind_group(0, bind_group, &[]);
            compute_pass.dispatch_workgroups((self.layouts.len() as u32).div_ceil(64), 1, 1);
        }
    }

    pub(super) fn update(
        &mut self,
        device: &wgpu::Device,
        queue: &wgpu::Queue,
        configs: &crate::wgpu::render_configs::RenderInputs,
        skinning_cache: &SkinningCache,
        bone_transforms_buffer: &wgpu::Buffer,
        bone_transforms_generation: u64,
    ) {
        let snapshot = &configs.scene.render_animation;

        if snapshot.signature != self.cached_static_signature
            || self.bone_transforms_generation != bone_transforms_generation
        {
            self.rebuild_static(device, snapshot, skinning_cache);
            self.cached_static_signature = snapshot.signature;
            self.bone_transforms_generation = bone_transforms_generation;
            self.bind_group = Some(self.create_bind_group(device, bone_transforms_buffer));
        }

        self.upload_chains(queue, snapshot);
    }

    fn rebuild_static(
        &mut self,
        device: &wgpu::Device,
        snapshot: &SkinnedAnimationSnapshot,
        skinning_cache: &SkinningCache,
    ) {
        self.layouts.clear();
        let mut joint_indices: Vec<u32> = Vec::new();
        let mut tip_descendants: Vec<u32> = Vec::new();
        let mut rest_axes: Vec<[f32; 4]> = Vec::new();
        let mut bone_lengths: Vec<f32> = Vec::new();
        let mut total_slots: u32 = 0;

        for skeleton in &snapshot.skeletons {
            if skeleton.spring_chains.is_empty() {
                continue;
            }
            let Some(&skin_index) = skinning_cache
                .entity_skin_indices
                .get(&skeleton.skin_entity)
            else {
                continue;
            };
            let base_bone_index = skinning_cache.get_base_bone_index(skin_index);

            let mut parent_of: std::collections::HashMap<u32, Option<u32>> =
                std::collections::HashMap::new();
            for joint in &skeleton.joints_ordered {
                parent_of.insert(joint.local_index, joint.parent_local);
            }

            for chain in &skeleton.spring_chains {
                if chain.joint_locals.len() < 2 {
                    continue;
                }
                let joint_start = joint_indices.len() as u32;
                for (index, &local) in chain.joint_locals.iter().enumerate() {
                    joint_indices.push(base_bone_index + local);
                    let axis = chain
                        .rest_axes
                        .get(index)
                        .copied()
                        .unwrap_or([0.0, 1.0, 0.0]);
                    rest_axes.push([axis[0], axis[1], axis[2], 0.0]);
                    bone_lengths.push(chain.bone_lengths.get(index).copied().unwrap_or(0.0));
                }
                let joint_count = chain.joint_locals.len() as u32;

                let tip_local = *chain.joint_locals.last().unwrap();
                let descendant_start = tip_descendants.len() as u32;
                for joint in &skeleton.joints_ordered {
                    if joint.local_index == tip_local {
                        continue;
                    }
                    let mut current = joint.parent_local;
                    let mut passes_tip = false;
                    while let Some(parent) = current {
                        if parent == tip_local {
                            passes_tip = true;
                            break;
                        }
                        current = parent_of.get(&parent).copied().flatten();
                    }
                    if passes_tip {
                        tip_descendants.push(base_bone_index + joint.local_index);
                    }
                }
                let descendant_count = tip_descendants.len() as u32 - descendant_start;

                self.layouts.push(SpringChainLayout {
                    player: skeleton.player_entity,
                    chain_index: chain.chain_index as usize,
                    joint_start,
                    joint_count,
                    descendant_start,
                    descendant_count,
                    state_start: total_slots,
                });
                total_slots += joint_count;
            }
        }

        self.joint_indices_buffer = init_buffer(
            device,
            "Spring Joint Indices",
            bytemuck::cast_slice(&joint_indices),
        );
        self.tip_descendants_buffer = init_buffer(
            device,
            "Spring Tip Descendants",
            bytemuck::cast_slice(&tip_descendants),
        );
        self.rest_axes_buffer =
            init_buffer(device, "Spring Rest Axes", bytemuck::cast_slice(&rest_axes));
        self.bone_lengths_buffer = init_buffer(
            device,
            "Spring Bone Lengths",
            bytemuck::cast_slice(&bone_lengths),
        );
        let state_bytes = (total_slots.max(1) as usize) * 2 * std::mem::size_of::<[f32; 4]>();
        self.state_buffer = init_buffer(device, "Spring State", &vec![0u8; state_bytes]);
        self.chains_buffer = sized_buffer(
            device,
            "Spring Chains",
            (self.layouts.len().max(1) * std::mem::size_of::<GpuSpringChain>()) as u64,
        );
        self.reset_pending = true;
    }

    fn upload_chains(&mut self, queue: &wgpu::Queue, snapshot: &SkinnedAnimationSnapshot) {
        if self.layouts.is_empty() {
            return;
        }
        let reset = u32::from(self.reset_pending);
        let chains: Vec<GpuSpringChain> = self
            .layouts
            .iter()
            .map(|layout| {
                let runtime = snapshot
                    .spring_runtime
                    .get(&(layout.player, layout.chain_index))
                    .copied()
                    .unwrap_or([0.0; 8]);
                GpuSpringChain {
                    joint_start: layout.joint_start,
                    joint_count: layout.joint_count,
                    descendant_start: layout.descendant_start,
                    descendant_count: layout.descendant_count,
                    state_start: layout.state_start,
                    reset,
                    enabled: u32::from(runtime[5] > 0.0),
                    pad0: 0,
                    gravity: [runtime[0], runtime[1], runtime[2], 0.0],
                    stiffness: runtime[3].clamp(0.0, 1.0),
                    damping: runtime[4].clamp(0.0, 1.0),
                    dt: runtime[6],
                    pad1: 0.0,
                }
            })
            .collect();
        queue.write_buffer(&self.chains_buffer, 0, bytemuck::cast_slice(&chains));
        self.reset_pending = false;
    }

    fn create_bind_group(
        &self,
        device: &wgpu::Device,
        bone_transforms_buffer: &wgpu::Buffer,
    ) -> wgpu::BindGroup {
        device.create_bind_group(&wgpu::BindGroupDescriptor {
            label: Some("Skinned Spring Bind Group"),
            layout: &self.bind_group_layout,
            entries: &[
                wgpu::BindGroupEntry {
                    binding: 0,
                    resource: bone_transforms_buffer.as_entire_binding(),
                },
                wgpu::BindGroupEntry {
                    binding: 1,
                    resource: self.chains_buffer.as_entire_binding(),
                },
                wgpu::BindGroupEntry {
                    binding: 2,
                    resource: self.joint_indices_buffer.as_entire_binding(),
                },
                wgpu::BindGroupEntry {
                    binding: 3,
                    resource: self.tip_descendants_buffer.as_entire_binding(),
                },
                wgpu::BindGroupEntry {
                    binding: 4,
                    resource: self.rest_axes_buffer.as_entire_binding(),
                },
                wgpu::BindGroupEntry {
                    binding: 5,
                    resource: self.bone_lengths_buffer.as_entire_binding(),
                },
                wgpu::BindGroupEntry {
                    binding: 6,
                    resource: self.state_buffer.as_entire_binding(),
                },
            ],
        })
    }
}