#![cfg(all(feature = "par", feature = "gpu-wgpu"))]
use alloc::format;
use alloc::string::ToString;
use alloc::vec::Vec;
use wgpu::{Buffer, Device, Queue};
use super::buffer::{create_download_buffer, download_buffer};
use super::error::{GpuError, GpuResult};
use super::kernel::KernelCache;
use crate::par::codegen::wgsl::{
KernelKey, OperatioReductionis, generate_map_shader, generate_reduce_shader,
};
pub(crate) const WORKGROUP_SIZE: usize = 64;
#[inline]
fn create_params_buffer(device: &Device) -> Buffer {
device.create_buffer(&wgpu::BufferDescriptor {
label: Some("parflumen_params"),
size: 4,
usage: wgpu::BufferUsages::UNIFORM | wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
})
}
#[inline]
fn to_u32(n: usize, what: &str) -> GpuResult<u32> {
u32::try_from(n).map_err(|_| {
GpuError::BufferCreationFailed(format!(
"{what} {n} exceeds u32::MAX — GPU path does not support this size"
))
})
}
#[inline]
pub(crate) fn execute_map_gpu_bytes(
device: &Device,
queue: &Queue,
cache: &mut KernelCache,
input_bytes: &[u8],
operation: &str,
type_name: &str,
element_size: usize,
) -> GpuResult<Vec<u8>> {
if input_bytes.is_empty() {
return Ok(Vec::new());
}
let (output_buffer, output_size_bytes) = execute_map_to_buffer(
device,
queue,
cache,
input_bytes,
operation,
type_name,
element_size,
)?;
let download_buf = create_download_buffer(device, output_size_bytes);
let mut encoder =
device.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
encoder.copy_buffer_to_buffer(
&output_buffer,
0,
&download_buf,
0,
output_size_bytes as u64,
);
let command_buffer = encoder.finish();
queue.submit([command_buffer]);
let result = download_buffer(
device,
queue,
&output_buffer,
&download_buf,
output_size_bytes,
);
cache.release_buffer(output_buffer);
result
}
#[inline]
pub(crate) fn execute_map_to_buffer(
device: &Device,
queue: &Queue,
cache: &mut KernelCache,
input_bytes: &[u8],
operation: &str,
type_name: &str,
element_size: usize,
) -> GpuResult<(Buffer, usize)> {
if input_bytes.is_empty() {
let buffer = cache.acquire_buffer(4); return Ok((buffer, 0));
}
let element_count = input_bytes.len() / element_size;
let element_count_u32 = to_u32(element_count, "element count")?;
let params_buffer = create_params_buffer(device);
queue.write_buffer(¶ms_buffer, 0, &element_count_u32.to_le_bytes());
let workgroup_size_u32 = to_u32(WORKGROUP_SIZE, "workgroup size")?;
let shader_source = generate_map_shader("map", operation, workgroup_size_u32, type_name);
let kernel_key = KernelKey::new(shader_source, "map".to_string());
let input_buffer = cache.acquire_buffer(input_bytes.len());
queue.write_buffer(&input_buffer, 0, input_bytes);
let output_buffer = cache.acquire_buffer(input_bytes.len());
let cached = match cache.get_or_compile(&kernel_key) {
Ok(c) => c,
Err(e) => {
cache.release_buffer(input_buffer);
cache.release_buffer(output_buffer);
return Err(e);
}
};
let bind_group = device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &cached.bind_group_layout,
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: input_buffer.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: output_buffer.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: params_buffer.as_entire_binding(),
},
],
});
let mut encoder =
device.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
{
let mut compute_pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: None,
timestamp_writes: None,
});
compute_pass.set_pipeline(&cached.pipeline);
compute_pass.set_bind_group(0, &bind_group, &[]);
let workgroup_count =
match to_u32(element_count.div_ceil(WORKGROUP_SIZE), "workgroup count") {
Ok(v) => v,
Err(e) => {
cache.release_buffer(input_buffer);
cache.release_buffer(output_buffer);
return Err(e);
}
};
compute_pass.dispatch_workgroups(workgroup_count, 1, 1);
}
let command_buffer = encoder.finish();
queue.submit([command_buffer]);
cache.release_buffer(input_buffer);
Ok((output_buffer, input_bytes.len()))
}
#[allow(clippy::too_many_lines)]
pub(crate) unsafe fn execute_reduce_from_buffer<T>(
device: &Device,
queue: &Queue,
cache: &mut KernelCache,
input_buffer: Buffer,
element_count: usize,
operation: &str,
type_name: &str,
) -> GpuResult<Option<T>>
where
T: Clone + Send + Sync + 'static,
{
if element_count == 0 {
cache.release_buffer(input_buffer);
return Ok(None);
}
let element_size = core::mem::size_of::<T>();
if element_count == 1 {
let download_buf = create_download_buffer(device, element_size);
let result_bytes =
match download_buffer(device, queue, &input_buffer, &download_buf, element_size) {
Ok(b) => b,
Err(e) => {
cache.release_buffer(input_buffer);
return Err(e);
}
};
cache.release_buffer(input_buffer);
if result_bytes.len() != core::mem::size_of::<T>() {
return Ok(None);
}
let t = unsafe { core::ptr::read_unaligned(result_bytes.as_ptr().cast::<T>()) };
return Ok(Some(t));
}
let workgroup_size = WORKGROUP_SIZE;
let op = match OperatioReductionis::parse(operation).ok_or_else(|| {
GpuError::ShaderCompilationFailed(format!(
"unsupported reduce operator `{operation}` (supported: + * min max)"
))
}) {
Ok(op) => op,
Err(e) => {
cache.release_buffer(input_buffer);
return Err(e);
}
};
let workgroup_size_u32 = match to_u32(workgroup_size, "workgroup size") {
Ok(v) => v,
Err(e) => {
cache.release_buffer(input_buffer);
return Err(e);
}
};
let shader_source = generate_reduce_shader("reduce", op, workgroup_size_u32, type_name);
let kernel_key = KernelKey::new(shader_source, "reduce".to_string());
if let Err(e) = cache.get_or_compile(&kernel_key) {
cache.release_buffer(input_buffer);
return Err(e);
}
let max_output_count = element_count.div_ceil(workgroup_size);
let scratch_size_bytes = max_output_count * element_size;
let mut scratch_a = cache.acquire_buffer(scratch_size_bytes);
let mut scratch_b = cache.acquire_buffer(scratch_size_bytes);
let params_buffer = create_params_buffer(device);
let mut current_count = element_count;
let mut first_pass = true;
while current_count > 1 {
let current_count_u32 = match to_u32(current_count, "element count") {
Ok(v) => v,
Err(e) => {
cache.release_buffer(input_buffer);
cache.release_buffer(scratch_a);
cache.release_buffer(scratch_b);
return Err(e);
}
};
queue.write_buffer(¶ms_buffer, 0, ¤t_count_u32.to_le_bytes());
let workgroup_count = current_count.div_ceil(workgroup_size);
let (input_res, output_res) = if first_pass {
(
input_buffer.as_entire_binding(),
scratch_a.as_entire_binding(),
)
} else {
(scratch_a.as_entire_binding(), scratch_b.as_entire_binding())
};
let cached = match cache.get_or_compile(&kernel_key) {
Ok(c) => c,
Err(e) => {
cache.release_buffer(input_buffer);
cache.release_buffer(scratch_a);
cache.release_buffer(scratch_b);
return Err(e);
}
};
let bind_group = device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &cached.bind_group_layout,
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: input_res,
},
wgpu::BindGroupEntry {
binding: 1,
resource: output_res,
},
wgpu::BindGroupEntry {
binding: 2,
resource: params_buffer.as_entire_binding(),
},
],
});
let mut encoder =
device.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
{
let mut compute_pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: None,
timestamp_writes: None,
});
compute_pass.set_pipeline(&cached.pipeline);
compute_pass.set_bind_group(0, &bind_group, &[]);
let workgroup_count_u32 = match to_u32(workgroup_count, "workgroup count") {
Ok(v) => v,
Err(e) => {
cache.release_buffer(input_buffer);
cache.release_buffer(scratch_a);
cache.release_buffer(scratch_b);
return Err(e);
}
};
compute_pass.dispatch_workgroups(workgroup_count_u32, 1, 1);
}
queue.submit([encoder.finish()]);
current_count = workgroup_count;
if first_pass {
first_pass = false;
} else {
core::mem::swap(&mut scratch_a, &mut scratch_b);
}
}
cache.release_buffer(scratch_b);
cache.release_buffer(input_buffer);
let current_buffer = scratch_a;
let download_buf = create_download_buffer(device, element_size);
let result_bytes =
match download_buffer(device, queue, ¤t_buffer, &download_buf, element_size) {
Ok(b) => b,
Err(e) => {
cache.release_buffer(current_buffer);
return Err(e);
}
};
cache.release_buffer(current_buffer);
if result_bytes.len() != element_size {
return Ok(None);
}
if result_bytes.len() != core::mem::size_of::<T>() {
return Ok(None);
}
let t = unsafe { core::ptr::read_unaligned(result_bytes.as_ptr().cast::<T>()) };
Ok(Some(t))
}