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(),
},
],
})
}
}