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 GpuCcdChain {
    joint_start: u32,
    joint_count: u32,
    descendant_start: u32,
    descendant_count: u32,
    goal: [f32; 4],
    weight: f32,
    max_angle: f32,
    iterations: u32,
    pad0: u32,
}

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

pub(super) struct SkinnedMultiIkGpu {
    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,
    layouts: Vec<CcdChainLayout>,
    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 SkinnedMultiIkGpu {
    pub(super) fn new(device: &wgpu::Device) -> Self {
        let shader = crate::wgpu::shader_compose::compile_wgsl(
            device,
            "ccd_compute.wgsl",
            include_str!("../../../shaders/ccd_compute.wgsl"),
        );

        let entries: Vec<wgpu::BindGroupLayoutEntry> = (0..4)
            .map(|binding| wgpu::BindGroupLayoutEntry {
                binding,
                visibility: wgpu::ShaderStages::COMPUTE,
                ty: wgpu::BindingType::Buffer {
                    ty: wgpu::BufferBindingType::Storage {
                        read_only: binding != 0,
                    },
                    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 CCD Bind Group Layout"),
            entries: &entries,
        });

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

        let pipeline = device.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
            label: Some("Skinned CCD 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, "CCD Chains", 16),
            joint_indices_buffer: sized_buffer(device, "CCD Joint Indices", 16),
            tip_descendants_buffer: sized_buffer(device, "CCD Tip Descendants", 16),
            layouts: Vec::new(),
            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();

        for skeleton in &snapshot.skeletons {
            if skeleton.ik_multi_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.ik_multi_chains {
                if chain.joint_locals.len() < 2 {
                    continue;
                }
                let joint_start = joint_indices.len() as u32;
                for &local in &chain.joint_locals {
                    joint_indices.push(base_bone_index + local);
                }
                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(CcdChainLayout {
                    player: skeleton.player_entity,
                    chain_index: chain.chain_index as usize,
                    joint_start,
                    joint_count,
                    descendant_start,
                    descendant_count,
                });
            }
        }

        self.joint_indices_buffer = init_buffer(
            device,
            "CCD Joint Indices",
            bytemuck::cast_slice(&joint_indices),
        );
        self.tip_descendants_buffer = init_buffer(
            device,
            "CCD Tip Descendants",
            bytemuck::cast_slice(&tip_descendants),
        );
        self.chains_buffer = sized_buffer(
            device,
            "CCD Chains",
            (self.layouts.len().max(1) * std::mem::size_of::<GpuCcdChain>()) as u64,
        );
    }

    fn upload_chains(&self, queue: &wgpu::Queue, snapshot: &SkinnedAnimationSnapshot) {
        if self.layouts.is_empty() {
            return;
        }
        let chains: Vec<GpuCcdChain> = self
            .layouts
            .iter()
            .map(|layout| {
                let runtime = snapshot
                    .ik_multi_runtime
                    .get(&(layout.player, layout.chain_index))
                    .copied()
                    .unwrap_or([0.0; 8]);
                let weight = if runtime[4] > 0.0 { runtime[3] } else { 0.0 };
                let iterations = (runtime[6].max(1.0)) as u32;
                let max_angle = if runtime[5] > 0.0 {
                    runtime[5]
                } else {
                    std::f32::consts::PI
                };
                GpuCcdChain {
                    joint_start: layout.joint_start,
                    joint_count: layout.joint_count,
                    descendant_start: layout.descendant_start,
                    descendant_count: layout.descendant_count,
                    goal: [runtime[0], runtime[1], runtime[2], 0.0],
                    weight,
                    max_angle,
                    iterations,
                    pad0: 0,
                }
            })
            .collect();
        queue.write_buffer(&self.chains_buffer, 0, bytemuck::cast_slice(&chains));
    }

    fn create_bind_group(
        &self,
        device: &wgpu::Device,
        bone_transforms_buffer: &wgpu::Buffer,
    ) -> wgpu::BindGroup {
        device.create_bind_group(&wgpu::BindGroupDescriptor {
            label: Some("Skinned CCD 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(),
                },
            ],
        })
    }
}