use crate::backend::{
Backend, BufferUsages, DescriptorType, DeviceValue, Dispatch, DispatchGrid, Encoder,
GpuBackendError, GpuTimestamp, MaybeSendSync, ShaderBinding,
};
use crate::shader::{BindGroupLayoutInfo, ShaderArgsError};
use async_channel::RecvError;
use bytemuck::{AnyBitPattern, NoUninit};
use regex::Regex;
use smallvec::SmallVec;
use std::borrow::Cow;
use std::ops::RangeBounds;
use std::sync::Arc;
use wgpu::util::{BufferInitDescriptor, DeviceExt};
use wgpu::wgt::CommandEncoderDescriptor;
use wgpu::{
Adapter, BindGroupLayout, BindGroupLayoutDescriptor, BindGroupLayoutEntry, BindingType, Buffer,
BufferAddress, BufferBindingType, BufferDescriptor, BufferSlice, BufferView, CommandEncoder,
ComputePass, ComputePassDescriptor, ComputePipeline, ComputePipelineDescriptor, Device,
ExperimentalFeatures, Instance, PipelineCompilationOptions, PipelineLayoutDescriptor,
PollError, Queue, ShaderModule, ShaderRuntimeChecks, ShaderStages,
};
fn shader_runtime_checks() -> ShaderRuntimeChecks {
ShaderRuntimeChecks {
force_loop_bounding: true,
..ShaderRuntimeChecks::unchecked()
}
}
#[derive(Clone, Copy)]
pub struct WebGpuBufferSlice<'a> {
pub(crate) inner: BufferSlice<'a>,
pub(crate) byte_len: u64,
}
impl<'a> WebGpuBufferSlice<'a> {
pub fn from_wgpu(buffer: &'a wgpu::Buffer) -> Self {
Self {
inner: buffer.slice(..),
byte_len: buffer.size(),
}
}
}
impl<'a> From<WebGpuBufferSlice<'a>> for wgpu::BindingResource<'a> {
fn from(slice: WebGpuBufferSlice<'a>) -> Self {
wgpu::BindingResource::try_from(slice.inner).expect("buffer slice can not be empty")
}
}
#[cfg(feature = "push_constants")]
#[derive(Clone)]
pub struct WebGpuModule {
pub module: ShaderModule,
pub bindings: Vec<BindGroupLayoutEntry>,
}
#[cfg(not(feature = "push_constants"))]
pub type WebGpuModule = ShaderModule;
#[derive(Clone)]
pub struct WebGpuFunction {
pub pipeline: ComputePipeline,
pub bind_group_layouts: Arc<Vec<BindGroupLayout>>,
}
pub struct WebGpuPass {
pub(crate) pass: ComputePass<'static>,
pub(crate) device: Device,
}
impl WebGpuPass {
pub fn begin_dispatch<'a>(&'a mut self, function: &'a WebGpuFunction) -> WebGpuDispatch<'a> {
WebGpuDispatch::new(&self.device, &mut self.pass, function)
}
}
pub struct WebGpuEncoder {
pub(crate) encoder: CommandEncoder,
pub(crate) device: Device,
}
#[derive(Clone)]
pub struct WebGpu {
_instance: Instance, _adapter: Adapter, device: Device,
queue: Queue,
hacks: Vec<(Regex, String)>,
pub force_buffer_copy_src: bool,
spirv_passthrough_enabled: bool,
timestamp_supported: bool,
}
impl WebGpu {
pub async fn default() -> anyhow::Result<Self> {
Self::new(wgpu::Features::default(), wgpu::Limits::default()).await
}
#[allow(unused_mut)] pub async fn new(
mut features: wgpu::Features,
mut limits: wgpu::Limits,
) -> anyhow::Result<Self> {
let instance = wgpu::Instance::default();
let adapter = instance
.request_adapter(&wgpu::RequestAdapterOptions {
power_preference: wgpu::PowerPreference::HighPerformance,
..Default::default()
})
.await
.map_err(|_| anyhow::anyhow!("Failed to initialize gpu adapter."))?;
#[cfg(feature = "push_constants")]
{
features = features | wgpu::Features::IMMEDIATES | wgpu::Features::SUBGROUP;
limits.max_immediate_size = 128;
}
#[cfg(feature = "subgroup_ops")]
{
features = features | wgpu::Features::SUBGROUP;
}
let timestamp_supported = adapter.features().contains(wgpu::Features::TIMESTAMP_QUERY);
if timestamp_supported {
features |= wgpu::Features::TIMESTAMP_QUERY;
}
let is_vulkan_backend = adapter.get_info().backend == wgpu::Backend::Vulkan;
let spirv_passthrough_enabled = is_vulkan_backend
&& adapter
.features()
.contains(wgpu::Features::PASSTHROUGH_SHADERS);
if spirv_passthrough_enabled {
features |= wgpu::Features::PASSTHROUGH_SHADERS;
}
let experimental_features = if spirv_passthrough_enabled {
unsafe { ExperimentalFeatures::enabled() }
} else {
ExperimentalFeatures::default()
};
let (device, queue) = adapter
.request_device(&wgpu::DeviceDescriptor {
label: None,
required_features: features,
required_limits: limits,
memory_hints: Default::default(),
trace: wgpu::Trace::Off,
experimental_features,
})
.await
.map_err(|e| anyhow::anyhow!("{:?}", e))?;
Ok(Self {
_instance: instance,
_adapter: adapter,
device,
queue,
force_buffer_copy_src: false,
hacks: vec![],
spirv_passthrough_enabled,
timestamp_supported,
})
}
pub fn from_device(instance: Instance, adapter: Adapter, device: Device, queue: Queue) -> Self {
let timestamp_supported = device.features().contains(wgpu::Features::TIMESTAMP_QUERY);
let is_vulkan_backend = adapter.get_info().backend == wgpu::Backend::Vulkan;
let spirv_passthrough_enabled = is_vulkan_backend
&& device
.features()
.contains(wgpu::Features::PASSTHROUGH_SHADERS);
Self {
_instance: instance,
_adapter: adapter,
device,
queue,
force_buffer_copy_src: false,
hacks: vec![],
spirv_passthrough_enabled,
timestamp_supported,
}
}
pub fn append_hack(&mut self, regex: Regex, replace_pattern: String) {
self.hacks.push((regex, replace_pattern));
}
pub fn device(&self) -> &Device {
&self.device
}
pub fn queue(&self) -> &Queue {
&self.queue
}
pub fn timestamp_supported(&self) -> bool {
self.timestamp_supported
}
#[cfg(feature = "push_constants")]
fn load_module_spirv(&self, spirv_bytes: &[u8]) -> Result<WebGpuModule, WebGpuBackendError> {
let source = wgpu::util::make_spirv(spirv_bytes);
let shader_module = self
.device
.create_shader_module(wgpu::ShaderModuleDescriptor {
label: None,
source,
});
Ok(WebGpuModule {
module: shader_module,
bindings: Vec::new(),
})
}
#[cfg(not(feature = "push_constants"))]
fn load_module_spirv(&self, spirv_bytes: &[u8]) -> Result<WebGpuModule, WebGpuBackendError> {
let source = wgpu::util::make_spirv(spirv_bytes);
let shader_module = unsafe {
self.device.create_shader_module_trusted(
wgpu::ShaderModuleDescriptor {
label: None,
source,
},
shader_runtime_checks(),
)
};
Ok(shader_module)
}
pub fn load_module_spirv_passthrough(
&self,
spirv_bytes: &[u8],
) -> Result<WebGpuModule, WebGpuBackendError> {
if !self.spirv_passthrough_enabled {
return self.load_module_spirv(spirv_bytes);
}
let spirv = wgpu::util::make_spirv_raw(spirv_bytes);
let shader_module = unsafe {
self.device
.create_shader_module_passthrough(wgpu::ShaderModuleDescriptorPassthrough {
spirv: Some(spirv),
..Default::default()
})
};
#[cfg(feature = "push_constants")]
return Ok(WebGpuModule {
module: shader_module,
bindings: Vec::new(),
});
#[cfg(not(feature = "push_constants"))]
Ok(shader_module)
}
}
#[derive(thiserror::Error, Debug)]
pub enum WebGpuBackendError {
#[error(transparent)]
ShaderArg(#[from] ShaderArgsError),
#[error(transparent)]
BytemuckPod(#[from] bytemuck::PodCastError),
#[error("Failed to read buffer from GPU: {0}")]
BufferRead(RecvError),
#[error(transparent)]
DevicePoll(#[from] PollError),
#[error(transparent)]
Recv(#[from] RecvError),
#[error(transparent)]
MapRange(#[from] wgpu::MapRangeError),
#[error("Failed to parse SPIR-V: {0}")]
SpirVParse(String),
#[error("Naga validation failed: {0}")]
NagaValidation(String),
#[error("Failed to write WGSL: {0}")]
WgslWrite(String),
}
impl Backend for WebGpu {
const NAME: &'static str = "webgpu";
const TARGET: super::CompileTarget = super::CompileTarget::Wgsl;
type Error = WebGpuBackendError;
type Buffer<T: DeviceValue> = Buffer;
type BufferSlice<'b, T: DeviceValue> = WebGpuBufferSlice<'b>;
type Encoder = WebGpuEncoder;
type Pass = WebGpuPass;
type Timestamps = WebGpuTimestamps;
type Module = WebGpuModule;
type Function = WebGpuFunction;
type Dispatch<'a> = WebGpuDispatch<'a>;
fn as_webgpu(&self) -> Option<&WebGpu> {
Some(self)
}
fn load_module(&self, data: &str) -> Result<Self::Module, Self::Error> {
let mut data = data.replace("enable f16;", "").replace("f16", "f32");
for (reg, replace) in &self.hacks {
data = reg.replace_all(&data, replace).to_string();
}
let shader_module = unsafe {
self.device.create_shader_module_trusted(
wgpu::ShaderModuleDescriptor {
label: None,
source: wgpu::ShaderSource::Wgsl(Cow::Borrowed(&data)),
},
shader_runtime_checks(),
)
};
#[cfg(feature = "push_constants")]
{
Ok(WebGpuModule {
module: shader_module,
bindings: Vec::new(),
})
}
#[cfg(not(feature = "push_constants"))]
Ok(shader_module)
}
fn load_module_bytes(&self, bytes: &[u8]) -> Result<Self::Module, Self::Error> {
if bytes.len() >= 4 {
let magic = u32::from_le_bytes([bytes[0], bytes[1], bytes[2], bytes[3]]);
if magic == 0x07230203 {
return self.load_module_spirv(bytes);
}
}
self.load_module(str::from_utf8(bytes).unwrap())
}
fn load_function(
&self,
module: &Self::Module,
entry_point: &str,
push_constant_size: u32,
) -> Result<Self::Function, Self::Error> {
self.load_function_with_layouts(
module,
entry_point,
push_constant_size,
&BindGroupLayoutInfo::default(),
)
}
fn load_function_with_layouts(
&self,
module: &Self::Module,
entry_point: &str,
push_constant_size: u32,
layouts: &BindGroupLayoutInfo,
) -> Result<Self::Function, Self::Error> {
let bind_group_layouts: Vec<BindGroupLayout> = layouts
.groups
.iter()
.enumerate()
.map(|(set_idx, bindings)| {
let entries: Vec<BindGroupLayoutEntry> = bindings
.iter()
.map(|binding| BindGroupLayoutEntry {
binding: binding.index,
visibility: ShaderStages::COMPUTE,
ty: match binding.descriptor_type {
DescriptorType::Uniform => BindingType::Buffer {
ty: BufferBindingType::Uniform,
has_dynamic_offset: false,
min_binding_size: None,
},
DescriptorType::Storage { read_only } => BindingType::Buffer {
ty: BufferBindingType::Storage { read_only },
has_dynamic_offset: false,
min_binding_size: None,
},
},
count: None,
})
.collect();
self.device
.create_bind_group_layout(&BindGroupLayoutDescriptor {
label: Some(&format!("{}:set{}", entry_point, set_idx)),
entries: &entries,
})
})
.collect();
let (shader_module, pipeline_layout) = if !bind_group_layouts.is_empty() {
let layout_refs: Vec<_> = bind_group_layouts.iter().map(Some).collect();
let layout = self
.device
.create_pipeline_layout(&PipelineLayoutDescriptor {
label: Some(entry_point),
bind_group_layouts: &layout_refs,
immediate_size: push_constant_size,
});
#[cfg(feature = "push_constants")]
let sm = &module.module;
#[cfg(not(feature = "push_constants"))]
let sm = module;
(sm, Some(layout))
} else {
#[cfg(feature = "push_constants")]
let sm = &module.module;
#[cfg(not(feature = "push_constants"))]
let sm = module;
let _ = push_constant_size;
(sm, None)
};
let pipeline = self
.device
.create_compute_pipeline(&ComputePipelineDescriptor {
label: Some(entry_point),
layout: pipeline_layout.as_ref(),
module: shader_module,
entry_point: Some(entry_point),
compilation_options: PipelineCompilationOptions {
zero_initialize_workgroup_memory: false,
..Default::default()
},
cache: None,
});
Ok(WebGpuFunction {
pipeline,
bind_group_layouts: Arc::new(bind_group_layouts),
})
}
fn begin_encoding(&self) -> Self::Encoder {
WebGpuEncoder {
encoder: self
.device
.create_command_encoder(&CommandEncoderDescriptor::default()),
device: self.device.clone(),
}
}
fn begin_dispatch<'a>(
&'a self,
pass: &'a mut Self::Pass,
function: &'a Self::Function,
) -> WebGpuDispatch<'a> {
pass.begin_dispatch(function)
}
fn submit(&self, encoder: Self::Encoder) -> Result<(), Self::Error> {
let _ = self.queue.submit(Some(encoder.encoder.finish()));
Ok(())
}
fn init_buffer<T: DeviceValue + NoUninit>(
&self,
data: &[T],
mut usage: BufferUsages,
) -> Result<Self::Buffer<T>, Self::Error> {
if self.force_buffer_copy_src && !usage.contains(BufferUsages::MAP_READ) {
usage |= BufferUsages::COPY_SRC;
}
Ok(self.device.create_buffer_init(&BufferInitDescriptor {
label: None,
contents: bytemuck::try_cast_slice(data)?,
usage: usage.into(),
}))
}
fn uninit_buffer<T: DeviceValue + NoUninit>(
&self,
len: usize,
mut usage: BufferUsages,
) -> Result<Self::Buffer<T>, Self::Error> {
if self.force_buffer_copy_src && !usage.contains(BufferUsages::MAP_READ) {
usage |= BufferUsages::COPY_SRC;
}
let bytes_len = std::mem::size_of::<T>() as u64 * len as u64;
Ok(self.device.create_buffer(&BufferDescriptor {
label: None,
size: bytes_len,
usage: usage.into(),
mapped_at_creation: false,
}))
}
fn write_buffer<T: DeviceValue + NoUninit>(
&self,
buffer: &mut Self::Buffer<T>,
offset: u64,
data: &[T],
) -> Result<(), Self::Error> {
let elt_sz = std::mem::size_of::<T>() as u64;
self.queue
.write_buffer(buffer, offset * elt_sz, bytemuck::cast_slice(data));
Ok(())
}
fn synchronize(&self) -> Result<(), Self::Error> {
self.device.poll(wgpu::PollType::wait_indefinitely())?;
Ok(())
}
fn poll(&self) {
let _ = self.device.poll(wgpu::PollType::Poll);
}
async fn read_buffer<T: MaybeSendSync + DeviceValue + AnyBitPattern>(
&self,
buffer: &Self::Buffer<T>,
out: &mut [T],
) -> Result<(), Self::Error> {
let data = read_bytes(&self.device, buffer).await?;
let out_bytes = core::mem::size_of_val(out);
let copy_len = data.len().min(out_bytes);
#[allow(dead_code)]
if false {
let _ = bytemuck::try_cast_slice::<_, T>(&data[..copy_len])?;
}
unsafe {
core::ptr::copy_nonoverlapping(data.as_ptr(), out.as_mut_ptr() as *mut u8, copy_len);
}
drop(data);
buffer.unmap();
Ok(())
}
async fn slow_read_buffer<T: MaybeSendSync + DeviceValue + AnyBitPattern>(
&self,
buffer: &Self::Buffer<T>,
out: &mut [T],
) -> Result<(), Self::Error> {
let bytes_len = buffer.size() as usize;
let staging =
self.uninit_buffer::<u8>(bytes_len, BufferUsages::MAP_READ | BufferUsages::COPY_DST)?;
let mut encoder = self.begin_encoding();
encoder
.encoder
.copy_buffer_to_buffer(buffer, 0, &staging, 0, bytes_len as u64);
self.submit(encoder)?;
self.read_buffer(&staging, out).await
}
}
impl Encoder<WebGpu> for WebGpuEncoder {
fn begin_pass(&mut self, label: &str, timestamps: Option<&mut WebGpuTimestamps>) -> WebGpuPass {
let mut desc = wgpu::ComputePassDescriptor {
label: (!label.is_empty()).then_some(label),
timestamp_writes: None,
};
if let Some(timestamps) = timestamps
&& let Some((begin_idx, end_idx)) = timestamps.alloc_timestamp_pair(label.to_string())
{
desc.timestamp_writes = Some(wgpu::ComputePassTimestampWrites {
query_set: ×tamps.query_set,
beginning_of_pass_write_index: Some(begin_idx),
end_of_pass_write_index: Some(end_idx),
});
}
WebGpuPass {
pass: self.encoder.begin_compute_pass(&desc).forget_lifetime(),
device: self.device.clone(),
}
}
fn copy_buffer_to_buffer<T: DeviceValue + NoUninit>(
&mut self,
source: &<WebGpu as Backend>::Buffer<T>,
source_offset: usize,
target: &mut <WebGpu as Backend>::Buffer<T>,
target_offset: usize,
copy_len: usize,
) -> Result<(), WebGpuBackendError> {
wgpu::CommandEncoder::copy_buffer_to_buffer(
&mut self.encoder,
source,
source_offset as BufferAddress * size_of::<T>() as BufferAddress,
target,
target_offset as BufferAddress * size_of::<T>() as BufferAddress,
copy_len as BufferAddress * size_of::<T>() as BufferAddress,
);
Ok(())
}
}
impl<'a> Dispatch<'a, WebGpu> for WebGpuDispatch<'a> {
#[cfg(feature = "push_constants")]
fn set_push_constants(&mut self, data: &[u8]) {
self.push_constants.clear();
self.push_constants.extend_from_slice(data);
}
fn launch<'b>(
self,
grid: impl Into<DispatchGrid<'b, WebGpu>>,
_block_dim: [u32; 3],
) -> Result<(), WebGpuBackendError> {
if !self.launchable {
return Ok(());
}
self.pass.set_pipeline(&self.pipeline);
#[cfg(feature = "push_constants")]
if !self.push_constants.is_empty() {
self.pass.set_immediates(0, &self.push_constants);
}
if !self.bind_group_layouts.is_empty() {
for (space, layout) in self.bind_group_layouts.iter().enumerate() {
let entries: SmallVec<[_; 10]> = self
.args
.iter()
.filter(|(binding, _)| binding.space == space as u32)
.map(|(binding, input)| wgpu::BindGroupEntry {
binding: binding.index,
resource: (*input).into(),
})
.collect();
let bind_group = self.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout,
entries: &entries,
});
self.pass.set_bind_group(space as u32, &bind_group, &[]);
}
} else {
let mut spaces: SmallVec<[u32; 4]> =
self.args.iter().map(|(binding, _)| binding.space).collect();
spaces.sort();
spaces.dedup();
for space in spaces {
let entries: SmallVec<[_; 10]> = self
.args
.iter()
.filter(|(binding, _)| binding.space == space)
.map(|(binding, input)| wgpu::BindGroupEntry {
binding: binding.index,
resource: (*input).into(),
})
.collect();
let pipeline = &self.pipeline;
let layout_result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
pipeline.get_bind_group_layout(space)
}));
if let Ok(layout) = layout_result {
let bind_group = self.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &layout,
entries: &entries,
});
self.pass.set_bind_group(space, &bind_group, &[]);
}
}
}
match grid.into() {
DispatchGrid::Grid(grid_dim) => {
if grid_dim[0] * grid_dim[1] * grid_dim[2] > 0 {
self.pass
.dispatch_workgroups(grid_dim[0], grid_dim[1], grid_dim[2]);
}
}
DispatchGrid::ThreadCount(threads) => {
let grid_dim = [
threads[0].div_ceil(_block_dim[0]),
threads[1].div_ceil(_block_dim[1]),
threads[2].div_ceil(_block_dim[2]),
];
if grid_dim[0] * grid_dim[1] * grid_dim[2] > 0 {
self.pass
.dispatch_workgroups(grid_dim[0], grid_dim[1], grid_dim[2]);
}
}
DispatchGrid::Indirect(grid_indirect) => {
self.pass.dispatch_workgroups_indirect(grid_indirect, 0);
}
}
Ok(())
}
}
pub struct WebGpuDispatch<'a> {
device: Device,
pass: &'a mut ComputePass<'static>,
pipeline: ComputePipeline,
bind_group_layouts: Arc<Vec<BindGroupLayout>>,
pub(crate) args: SmallVec<[(ShaderBinding, WebGpuBufferSlice<'a>); 10]>,
launchable: bool,
#[cfg(feature = "push_constants")]
push_constants: SmallVec<[u8; 128]>,
}
impl<'a> WebGpuDispatch<'a> {
fn new(
device: &Device,
pass: &'a mut ComputePass<'static>,
function: &WebGpuFunction,
) -> WebGpuDispatch<'a> {
WebGpuDispatch {
device: device.clone(),
pass,
pipeline: function.pipeline.clone(),
bind_group_layouts: function.bind_group_layouts.clone(),
args: SmallVec::default(),
launchable: true,
#[cfg(feature = "push_constants")]
push_constants: SmallVec::default(),
}
}
}
#[allow(dead_code)]
pub trait CommandEncoderExt {
fn compute_pass<'encoder>(&'encoder mut self, label: &str) -> ComputePass<'encoder>;
}
impl CommandEncoderExt for CommandEncoder {
fn compute_pass<'encoder>(&'encoder mut self, label: &str) -> ComputePass<'encoder> {
let desc = ComputePassDescriptor {
label: Some(label),
timestamp_writes: None,
};
self.begin_compute_pass(&desc)
}
}
impl CommandEncoderExt for WebGpuEncoder {
fn compute_pass<'encoder>(&'encoder mut self, label: &str) -> ComputePass<'encoder> {
self.encoder.compute_pass(label)
}
}
async fn read_bytes(device: &Device, buffer: &Buffer) -> Result<BufferView, WebGpuBackendError> {
let buffer_slice = buffer.slice(..);
#[cfg(not(target_arch = "wasm32"))]
{
let (sender, receiver) = async_channel::bounded(1);
buffer_slice.map_async(wgpu::MapMode::Read, move |v| {
sender.send_blocking(v).unwrap()
});
device.poll(wgpu::PollType::wait_indefinitely())?;
receiver
.recv()
.await
.map_err(WebGpuBackendError::BufferRead)?
.unwrap();
}
#[cfg(target_arch = "wasm32")]
{
let (sender, receiver) = async_channel::bounded(1);
buffer_slice.map_async(wgpu::MapMode::Read, move |v| {
let _ = sender.force_send(v).unwrap();
});
device.poll(wgpu::PollType::wait_indefinitely())?;
receiver.recv().await?.unwrap();
}
let data = buffer_slice.get_mapped_range()?;
Ok(data)
}
impl<T: DeviceValue> crate::backend::Buffer<WebGpu, T> for Buffer {
fn is_empty(&self) -> bool {
self.size() == 0
}
fn len(&self) -> usize
where
T: Sized,
{
self.size() as usize / std::mem::size_of::<T>()
}
fn slice(&self, range: impl RangeBounds<usize>) -> <WebGpu as Backend>::BufferSlice<'_, T> {
let elem_size = std::mem::size_of::<T>() as u64;
let start_bytes = match range.start_bound() {
std::ops::Bound::Included(&val) => val as u64 * elem_size,
std::ops::Bound::Unbounded => 0,
_ => unreachable!(),
};
let end_bytes = match range.end_bound() {
std::ops::Bound::Excluded(&val) => val as u64 * elem_size,
std::ops::Bound::Included(&val) => (val as u64 + 1) * elem_size,
std::ops::Bound::Unbounded => self.size(),
};
WebGpuBufferSlice {
inner: self.slice(start_bytes..end_bytes),
byte_len: end_bytes - start_bytes,
}
}
fn usage(&self) -> BufferUsages {
self.usage().into()
}
}
enum TimestampReadState {
Idle,
Empty,
Pending {
rx: async_channel::Receiver<Result<(), wgpu::BufferAsyncError>>,
query_count: u32,
},
}
pub struct WebGpuTimestamps {
query_set: wgpu::QuerySet,
resolve_buffer: wgpu::Buffer,
staging_buffer: wgpu::Buffer,
capacity: u32,
next_index: u32,
labels: Vec<String>,
timestamp_period: f32,
read_state: TimestampReadState,
}
impl WebGpuTimestamps {
pub fn new(backend: &WebGpu, capacity: u32) -> Option<Self> {
if !backend.timestamp_supported {
return None;
}
let query_count = capacity * 2; let query_set = backend.device.create_query_set(&wgpu::QuerySetDescriptor {
label: Some("gpu_timestamps"),
ty: wgpu::QueryType::Timestamp,
count: query_count,
});
let resolve_buffer_size = (query_count as u64) * std::mem::size_of::<u64>() as u64;
let resolve_buffer = backend.device.create_buffer(&wgpu::BufferDescriptor {
label: Some("gpu_timestamps_resolve"),
size: resolve_buffer_size,
usage: wgpu::BufferUsages::QUERY_RESOLVE | wgpu::BufferUsages::COPY_SRC,
mapped_at_creation: false,
});
let staging_buffer = backend.device.create_buffer(&wgpu::BufferDescriptor {
label: Some("gpu_timestamps_staging"),
size: resolve_buffer_size,
usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
});
let timestamp_period = backend.queue.get_timestamp_period();
Some(WebGpuTimestamps {
query_set,
resolve_buffer,
staging_buffer,
capacity,
next_index: 0,
labels: Vec::with_capacity(capacity as usize),
timestamp_period,
read_state: TimestampReadState::Idle,
})
}
pub fn reset(&mut self) {
self.next_index = 0;
self.labels.clear();
self.read_state = TimestampReadState::Idle;
}
pub fn is_idle(&self) -> bool {
matches!(self.read_state, TimestampReadState::Idle)
}
pub fn resolve(&self, encoder: &mut WebGpuEncoder) {
if self.next_index == 0 {
return;
}
let query_count = self.next_index * 2;
encoder
.encoder
.resolve_query_set(&self.query_set, 0..query_count, &self.resolve_buffer, 0);
let copy_size = query_count as u64 * std::mem::size_of::<u64>() as u64;
encoder.encoder.copy_buffer_to_buffer(
&self.resolve_buffer,
0,
&self.staging_buffer,
0,
copy_size,
);
}
pub async fn read(&self, backend: &WebGpu) -> Result<Vec<GpuTimestamp>, GpuBackendError> {
if self.next_index == 0 {
return Ok(Vec::new());
}
let query_count = self.next_index * 2;
let mut raw_timestamps = vec![0u64; query_count as usize];
{
use crate::backend::webgpu::WebGpuBackendError;
let buffer_slice = self.staging_buffer.slice(..);
let (sender, receiver) = async_channel::bounded(1);
buffer_slice.map_async(wgpu::MapMode::Read, move |v| {
let _ = sender.force_send(v).unwrap();
});
backend
.device()
.poll(wgpu::PollType::wait_indefinitely())
.map_err(WebGpuBackendError::from)?;
receiver
.recv()
.await
.map_err(WebGpuBackendError::from)?
.unwrap();
let data = buffer_slice
.get_mapped_range()
.map_err(WebGpuBackendError::from)?;
let bytes = &*data;
unsafe {
std::ptr::copy_nonoverlapping(
bytes.as_ptr(),
raw_timestamps.as_mut_ptr() as *mut u8,
bytes
.len()
.min(raw_timestamps.len() * std::mem::size_of::<u64>()),
);
}
drop(data);
self.staging_buffer.unmap();
}
Ok(self.durations_from_raw(&raw_timestamps))
}
pub fn request_read(&mut self) {
if self.next_index == 0 {
self.read_state = TimestampReadState::Empty;
return;
}
let query_count = self.next_index * 2;
let (sender, receiver) = async_channel::bounded(1);
self.staging_buffer
.slice(..)
.map_async(wgpu::MapMode::Read, move |v| {
let _ = sender.force_send(v);
});
self.read_state = TimestampReadState::Pending {
rx: receiver,
query_count,
};
}
pub fn try_take(&mut self) -> Option<Vec<GpuTimestamp>> {
match std::mem::replace(&mut self.read_state, TimestampReadState::Idle) {
TimestampReadState::Idle => None,
TimestampReadState::Empty => Some(Vec::new()),
TimestampReadState::Pending { rx, query_count } => match rx.try_recv() {
Ok(Ok(())) => {
let mut raw = vec![0u64; query_count as usize];
let buffer_slice = self.staging_buffer.slice(..);
let Ok(data) = buffer_slice.get_mapped_range() else {
self.staging_buffer.unmap();
return Some(Vec::new());
};
unsafe {
std::ptr::copy_nonoverlapping(
data.as_ptr(),
raw.as_mut_ptr() as *mut u8,
data.len().min(raw.len() * std::mem::size_of::<u64>()),
);
}
drop(data);
self.staging_buffer.unmap();
Some(self.durations_from_raw(&raw))
}
Ok(Err(_)) | Err(async_channel::TryRecvError::Closed) => Some(Vec::new()),
Err(async_channel::TryRecvError::Empty) => {
self.read_state = TimestampReadState::Pending { rx, query_count };
None
}
},
}
}
fn durations_from_raw(&self, raw: &[u64]) -> Vec<GpuTimestamp> {
let period_ms = self.timestamp_period as f64 / 1_000_000.0;
self.labels
.iter()
.enumerate()
.map(|(i, label)| {
let begin = raw[i * 2];
let end = raw[i * 2 + 1];
let duration_ms = (end.wrapping_sub(begin)) as f64 * period_ms;
GpuTimestamp {
label: label.clone(),
duration_ms,
}
})
.collect()
}
pub fn alloc_timestamp_pair(&mut self, label: String) -> Option<(u32, u32)> {
if self.next_index < self.capacity {
let begin_idx = self.next_index * 2;
let end_idx = begin_idx + 1;
self.next_index += 1;
self.labels.push(label);
Some((begin_idx, end_idx))
} else {
None
}
}
}