use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use oxicuda_backend::{
BackendError, BackendResult, BackendTranspose, BinaryOp, ComputeBackend, ReduceOp, UnaryOp,
};
use wgpu;
use crate::{
device::WebGpuDevice,
memory::WebGpuMemoryManager,
planner::{self, Limits},
shader,
};
#[path = "backend_gpu_ops.rs"]
mod gpu_ops;
use gpu_ops::{
attention_cpu_reference, attention_gpu_dispatch_grid, conv2d_cpu_reference,
conv2d_gpu_dispatch_grid, conv2d_u32_dims,
};
#[path = "backend_cache.rs"]
mod cache;
use cache::{BindGroupCache, CachedPipeline};
fn map_unary_op(op: UnaryOp) -> &'static str {
match op {
UnaryOp::Relu => "relu",
UnaryOp::Sigmoid => "sigmoid",
UnaryOp::Tanh => "tanh",
UnaryOp::Exp => "exp",
UnaryOp::Log => "log",
UnaryOp::Sqrt => "sqrt",
UnaryOp::Abs => "abs",
UnaryOp::Neg => "neg",
}
}
fn map_binary_op(op: BinaryOp) -> &'static str {
match op {
BinaryOp::Add => "add",
BinaryOp::Sub => "sub",
BinaryOp::Mul => "mul",
BinaryOp::Div => "div",
BinaryOp::Max => "max",
BinaryOp::Min => "min",
}
}
fn map_reduce_op(op: ReduceOp) -> &'static str {
match op {
ReduceOp::Sum => "sum",
ReduceOp::Max => "max",
ReduceOp::Min => "min",
ReduceOp::Mean => "mean",
}
}
fn packed_gemm_lds(
trans_a: BackendTranspose,
trans_b: BackendTranspose,
m: usize,
n: usize,
k: usize,
) -> (usize, usize, usize) {
let lda = if trans_a == BackendTranspose::NoTrans {
k
} else {
m
};
let ldb = if trans_b == BackendTranspose::NoTrans {
n
} else {
k
};
(lda, ldb, n)
}
fn dim_u32(context: &str, name: &str, value: usize) -> BackendResult<u32> {
u32::try_from(value).map_err(|_| {
BackendError::InvalidArgument(format!("{context}: {name} {value} exceeds u32 range"))
})
}
fn gpu_limits() -> Limits {
Limits::portable_default()
}
#[derive(Debug)]
pub struct WebGpuBackend {
device: Option<Arc<WebGpuDevice>>,
memory: Option<Arc<WebGpuMemoryManager>>,
initialized: bool,
pipeline_cache: Mutex<HashMap<String, CachedPipeline>>,
bind_group_cache: Mutex<BindGroupCache>,
}
impl WebGpuBackend {
pub fn new() -> Self {
Self {
device: None,
memory: None,
initialized: false,
pipeline_cache: Mutex::new(HashMap::new()),
bind_group_cache: Mutex::new(BindGroupCache::new()),
}
}
fn check_init(&self) -> BackendResult<()> {
if self.initialized {
Ok(())
} else {
Err(BackendError::NotInitialized)
}
}
fn memory(&self) -> BackendResult<&Arc<WebGpuMemoryManager>> {
self.memory.as_ref().ok_or(BackendError::NotInitialized)
}
fn device(&self) -> BackendResult<&Arc<WebGpuDevice>> {
self.device.as_ref().ok_or(BackendError::NotInitialized)
}
#[must_use]
pub fn supports_f16(&self) -> bool {
self.device.as_ref().is_some_and(|d| d.supports_f16)
}
fn reduce_nd(
&self,
op: ReduceOp,
input_ptr: u64,
output_ptr: u64,
shape: &[usize],
axis: usize,
) -> BackendResult<()> {
debug_assert!(!shape.is_empty());
debug_assert!(axis < shape.len());
let outer: usize = shape[..axis].iter().product();
let dk: usize = shape[axis];
let inner: usize = shape[axis + 1..].iter().product();
if outer == 0 || dk == 0 || inner == 0 {
return Ok(());
}
let total = outer.checked_mul(inner).ok_or_else(|| {
BackendError::InvalidArgument("reduce: outer * inner overflows usize".into())
})?;
let in_elems = outer
.checked_mul(dk)
.and_then(|v| v.checked_mul(inner))
.ok_or_else(|| {
BackendError::InvalidArgument("reduce: outer * dk * inner overflows usize".into())
})?;
let inner_stride: usize = 1;
let dk_stride: usize = inner;
let outer_stride: usize = dk
.checked_mul(inner)
.ok_or_else(|| BackendError::InvalidArgument("reduce: dk * inner overflows".into()))?;
let limits = gpu_limits();
let (grid, grid_x) = planner::plan_dispatch_1d(&limits, total as u64, 1)
.map_err(BackendError::InvalidArgument)?;
let dev = self.device()?;
let mem = self.memory()?;
let op_str = map_reduce_op(op);
let pipeline_key = format!("reduce_nd:{op_str}");
let cached = self.cached_pipeline(&pipeline_key, "oxicuda-reduce-nd", || {
shader::reduction_nd_wgsl(op_str)
})?;
let mut params_bytes = [0u8; 32];
let outer_u32: u32 = outer
.try_into()
.map_err(|_| BackendError::InvalidArgument("reduce: outer exceeds u32 range".into()))?;
let dk_u32: u32 = dk
.try_into()
.map_err(|_| BackendError::InvalidArgument("reduce: dk exceeds u32 range".into()))?;
let inner_u32: u32 = inner
.try_into()
.map_err(|_| BackendError::InvalidArgument("reduce: inner exceeds u32 range".into()))?;
let outer_stride_u32: u32 = outer_stride.try_into().map_err(|_| {
BackendError::InvalidArgument("reduce: outer_stride exceeds u32 range".into())
})?;
let dk_stride_u32: u32 = dk_stride.try_into().map_err(|_| {
BackendError::InvalidArgument("reduce: dk_stride exceeds u32 range".into())
})?;
let inner_stride_u32: u32 = inner_stride.try_into().map_err(|_| {
BackendError::InvalidArgument("reduce: inner_stride exceeds u32 range".into())
})?;
params_bytes[0..4].copy_from_slice(&outer_u32.to_le_bytes());
params_bytes[4..8].copy_from_slice(&dk_u32.to_le_bytes());
params_bytes[8..12].copy_from_slice(&inner_u32.to_le_bytes());
params_bytes[12..16].copy_from_slice(&outer_stride_u32.to_le_bytes());
params_bytes[16..20].copy_from_slice(&dk_stride_u32.to_le_bytes());
params_bytes[20..24].copy_from_slice(&inner_stride_u32.to_le_bytes());
params_bytes[24..28].copy_from_slice(&grid_x.to_le_bytes());
let need_in = (in_elems as u64) * 4;
let need_out = (total as u64) * 4;
let bind_group = self.cached_bind_group(
dev,
mem,
&cached.bind_group_layout,
&pipeline_key,
&[input_ptr, output_ptr],
&[need_in, need_out],
¶ms_bytes,
"oxicuda-reduce-nd",
)?;
let mut encoder = dev
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("oxicuda-reduce-nd"),
});
{
let mut pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: Some("oxicuda-reduce-nd"),
timestamp_writes: None,
});
pass.set_pipeline(&cached.pipeline);
pass.set_bind_group(0, &bind_group, &[]);
pass.dispatch_workgroups(grid.x, grid.y, grid.z);
}
dev.queue.submit(std::iter::once(encoder.finish()));
Ok(())
}
}
impl WebGpuBackend {
#[allow(clippy::too_many_arguments)]
pub fn gemm_f16(
&self,
trans_a: BackendTranspose,
trans_b: BackendTranspose,
m: usize,
n: usize,
k: usize,
alpha: f64,
a_ptr: u64,
lda: usize,
b_ptr: u64,
ldb: usize,
beta: f64,
c_ptr: u64,
ldc: usize,
) -> BackendResult<()> {
self.check_init()?;
if m == 0 || n == 0 || k == 0 {
return Ok(());
}
let dev = self.device()?;
let mem = self.memory()?;
if !dev.supports_f16 {
return Err(BackendError::Unsupported(
"f16 GEMM requires the SHADER_F16 device feature, \
which this adapter does not support"
.into(),
));
}
let trans_a_flag: u32 = u32::from(trans_a != BackendTranspose::NoTrans);
let trans_b_flag: u32 = u32::from(trans_b != BackendTranspose::NoTrans);
let (expected_lda, expected_ldb, expected_ldc) = packed_gemm_lds(trans_a, trans_b, m, n, k);
if lda < expected_lda || ldb < expected_ldb || ldc < expected_ldc {
return Err(BackendError::InvalidArgument(
"gemm_f16: leading dimension smaller than matrix extent".into(),
));
}
let m_u32 = dim_u32("gemm_f16", "m", m)?;
let n_u32 = dim_u32("gemm_f16", "n", n)?;
let k_u32 = dim_u32("gemm_f16", "k", k)?;
let lda_u32 = dim_u32("gemm_f16", "lda", lda)?;
let ldb_u32 = dim_u32("gemm_f16", "ldb", ldb)?;
let ldc_u32 = dim_u32("gemm_f16", "ldc", ldc)?;
let limits = gpu_limits();
let tile = planner::plan_workgroup_square(&limits, 16);
let tile_size = tile.x;
let grid = planner::plan_dispatch_2d(&limits, m_u32, n_u32, tile, 1)
.map_err(BackendError::InvalidArgument)?;
let pipeline_key = format!("gemm_f16:{tile_size}");
let cached = self.cached_pipeline(&pipeline_key, "oxicuda-gemm-f16", || {
shader::gemm_wgsl_f16(tile_size)
})?;
let mut params_bytes = [0u8; 48];
params_bytes[0..4].copy_from_slice(&m_u32.to_le_bytes());
params_bytes[4..8].copy_from_slice(&n_u32.to_le_bytes());
params_bytes[8..12].copy_from_slice(&k_u32.to_le_bytes());
params_bytes[12..16].copy_from_slice(&(alpha as f32).to_le_bytes());
params_bytes[16..20].copy_from_slice(&(beta as f32).to_le_bytes());
params_bytes[20..24].copy_from_slice(&trans_a_flag.to_le_bytes());
params_bytes[24..28].copy_from_slice(&trans_b_flag.to_le_bytes());
params_bytes[28..32].copy_from_slice(&lda_u32.to_le_bytes());
params_bytes[32..36].copy_from_slice(&ldb_u32.to_le_bytes());
params_bytes[36..40].copy_from_slice(&ldc_u32.to_le_bytes());
let bind_group = self.cached_bind_group(
dev,
mem,
&cached.bind_group_layout,
&pipeline_key,
&[a_ptr, b_ptr, c_ptr],
&[],
¶ms_bytes,
"oxicuda-gemm-f16",
)?;
let mut encoder = dev
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("oxicuda-gemm-f16"),
});
{
let mut pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: Some("oxicuda-gemm-f16"),
timestamp_writes: None,
});
pass.set_pipeline(&cached.pipeline);
pass.set_bind_group(0, &bind_group, &[]);
pass.dispatch_workgroups(grid.x, grid.y, grid.z);
}
dev.queue.submit(std::iter::once(encoder.finish()));
Ok(())
}
}
impl Default for WebGpuBackend {
fn default() -> Self {
Self::new()
}
}
impl ComputeBackend for WebGpuBackend {
fn name(&self) -> &str {
"webgpu"
}
fn init(&mut self) -> BackendResult<()> {
if self.initialized {
return Ok(());
}
match WebGpuDevice::new() {
Ok(dev) => {
let dev = Arc::new(dev);
tracing::info!("WebGPU backend initialised on: {}", dev.adapter_name);
let memory = WebGpuMemoryManager::new(Arc::clone(&dev));
self.device = Some(dev);
self.memory = Some(Arc::new(memory));
self.initialized = true;
Ok(())
}
Err(e) => Err(BackendError::from(e)),
}
}
fn is_initialized(&self) -> bool {
self.initialized
}
fn gemm(
&self,
trans_a: BackendTranspose,
trans_b: BackendTranspose,
m: usize,
n: usize,
k: usize,
alpha: f64,
a_ptr: u64,
lda: usize,
b_ptr: u64,
ldb: usize,
beta: f64,
c_ptr: u64,
ldc: usize,
) -> BackendResult<()> {
self.check_init()?;
if m == 0 || n == 0 || k == 0 {
return Ok(());
}
let trans_a_flag: u32 = u32::from(trans_a != BackendTranspose::NoTrans);
let trans_b_flag: u32 = u32::from(trans_b != BackendTranspose::NoTrans);
let dev = self.device()?;
let mem = self.memory()?;
let m_u32 = dim_u32("gemm", "m", m)?;
let n_u32 = dim_u32("gemm", "n", n)?;
let k_u32 = dim_u32("gemm", "k", k)?;
let limits = gpu_limits();
let tile = planner::plan_workgroup_square(&limits, 16);
let tile_size = tile.x;
let grid = planner::plan_dispatch_2d(&limits, m_u32, n_u32, tile, 1)
.map_err(BackendError::InvalidArgument)?;
let pipeline_key = format!("gemm:{tile_size}");
let cached = self.cached_pipeline(&pipeline_key, "oxicuda-gemm", || {
shader::gemm_wgsl(tile_size)
})?;
let (expected_lda, expected_ldb, expected_ldc) = packed_gemm_lds(trans_a, trans_b, m, n, k);
if lda < expected_lda || ldb < expected_ldb || ldc < expected_ldc {
return Err(BackendError::InvalidArgument(
"gemm: leading dimension smaller than matrix extent".into(),
));
}
let lda_u32 = dim_u32("gemm", "lda", lda)?;
let ldb_u32 = dim_u32("gemm", "ldb", ldb)?;
let ldc_u32 = dim_u32("gemm", "ldc", ldc)?;
let mut params_bytes = [0u8; 48];
params_bytes[0..4].copy_from_slice(&m_u32.to_le_bytes());
params_bytes[4..8].copy_from_slice(&n_u32.to_le_bytes());
params_bytes[8..12].copy_from_slice(&k_u32.to_le_bytes());
params_bytes[12..16].copy_from_slice(&(alpha as f32).to_le_bytes());
params_bytes[16..20].copy_from_slice(&(beta as f32).to_le_bytes());
params_bytes[20..24].copy_from_slice(&trans_a_flag.to_le_bytes());
params_bytes[24..28].copy_from_slice(&trans_b_flag.to_le_bytes());
params_bytes[28..32].copy_from_slice(&lda_u32.to_le_bytes());
params_bytes[32..36].copy_from_slice(&ldb_u32.to_le_bytes());
params_bytes[36..40].copy_from_slice(&ldc_u32.to_le_bytes());
let bind_group = self.cached_bind_group(
dev,
mem,
&cached.bind_group_layout,
&pipeline_key,
&[a_ptr, b_ptr, c_ptr],
&[],
¶ms_bytes,
"oxicuda-gemm",
)?;
let mut encoder = dev
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("oxicuda-gemm"),
});
{
let mut pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: Some("oxicuda-gemm"),
timestamp_writes: None,
});
pass.set_pipeline(&cached.pipeline);
pass.set_bind_group(0, &bind_group, &[]);
pass.dispatch_workgroups(grid.x, grid.y, grid.z);
}
dev.queue.submit(std::iter::once(encoder.finish()));
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn batched_gemm(
&self,
trans_a: BackendTranspose,
trans_b: BackendTranspose,
m: usize,
n: usize,
k: usize,
alpha: f64,
a_ptr: u64,
lda: usize,
stride_a: usize,
b_ptr: u64,
ldb: usize,
stride_b: usize,
beta: f64,
c_ptr: u64,
ldc: usize,
stride_c: usize,
batch_count: usize,
) -> BackendResult<()> {
self.check_init()?;
if batch_count == 0 || m == 0 || n == 0 || k == 0 {
return Ok(());
}
let trans_a_flag: u32 = u32::from(trans_a != BackendTranspose::NoTrans);
let trans_b_flag: u32 = u32::from(trans_b != BackendTranspose::NoTrans);
let dev = self.device()?;
let mem = self.memory()?;
let m_u32 = dim_u32("batched_gemm", "m", m)?;
let n_u32 = dim_u32("batched_gemm", "n", n)?;
let k_u32 = dim_u32("batched_gemm", "k", k)?;
let batch_u32 = dim_u32("batched_gemm", "batch_count", batch_count)?;
let stride_a_u32 = dim_u32("batched_gemm", "stride_a", stride_a)?;
let stride_b_u32 = dim_u32("batched_gemm", "stride_b", stride_b)?;
let stride_c_u32 = dim_u32("batched_gemm", "stride_c", stride_c)?;
let limits = gpu_limits();
let tile = planner::plan_workgroup_square(&limits, 16);
let tile_size = tile.x;
let grid = planner::plan_dispatch_2d(&limits, m_u32, n_u32, tile, batch_u32)
.map_err(BackendError::InvalidArgument)?;
let pipeline_key = format!("batched_gemm:{tile_size}");
let cached = self.cached_pipeline(&pipeline_key, "oxicuda-batched-gemm", || {
shader::batched_gemm_wgsl(tile_size)
})?;
let (expected_lda, expected_ldb, expected_ldc) = packed_gemm_lds(trans_a, trans_b, m, n, k);
if lda < expected_lda || ldb < expected_ldb || ldc < expected_ldc {
return Err(BackendError::InvalidArgument(
"batched_gemm: leading dimension smaller than matrix extent".into(),
));
}
let lda_u32 = dim_u32("batched_gemm", "lda", lda)?;
let ldb_u32 = dim_u32("batched_gemm", "ldb", ldb)?;
let ldc_u32 = dim_u32("batched_gemm", "ldc", ldc)?;
let mut params_bytes = [0u8; 64];
params_bytes[0..4].copy_from_slice(&m_u32.to_le_bytes());
params_bytes[4..8].copy_from_slice(&n_u32.to_le_bytes());
params_bytes[8..12].copy_from_slice(&k_u32.to_le_bytes());
params_bytes[12..16].copy_from_slice(&(alpha as f32).to_le_bytes());
params_bytes[16..20].copy_from_slice(&(beta as f32).to_le_bytes());
params_bytes[20..24].copy_from_slice(&batch_u32.to_le_bytes());
params_bytes[24..28].copy_from_slice(&stride_a_u32.to_le_bytes());
params_bytes[28..32].copy_from_slice(&stride_b_u32.to_le_bytes());
params_bytes[32..36].copy_from_slice(&stride_c_u32.to_le_bytes());
params_bytes[36..40].copy_from_slice(&trans_a_flag.to_le_bytes());
params_bytes[40..44].copy_from_slice(&trans_b_flag.to_le_bytes());
params_bytes[44..48].copy_from_slice(&lda_u32.to_le_bytes());
params_bytes[48..52].copy_from_slice(&ldb_u32.to_le_bytes());
params_bytes[52..56].copy_from_slice(&ldc_u32.to_le_bytes());
let bind_group = self.cached_bind_group(
dev,
mem,
&cached.bind_group_layout,
&pipeline_key,
&[a_ptr, b_ptr, c_ptr],
&[],
¶ms_bytes,
"oxicuda-batched-gemm",
)?;
let mut encoder = dev
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("oxicuda-batched-gemm"),
});
{
let mut pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: Some("oxicuda-batched-gemm"),
timestamp_writes: None,
});
pass.set_pipeline(&cached.pipeline);
pass.set_bind_group(0, &bind_group, &[]);
pass.dispatch_workgroups(grid.x, grid.y, grid.z);
}
dev.queue.submit(std::iter::once(encoder.finish()));
Ok(())
}
fn conv2d_forward(
&self,
input_ptr: u64,
input_shape: &[usize],
filter_ptr: u64,
filter_shape: &[usize],
output_ptr: u64,
output_shape: &[usize],
stride: &[usize],
padding: &[usize],
) -> BackendResult<()> {
self.check_init()?;
if input_shape.len() != 4 {
return Err(BackendError::InvalidArgument(
"input_shape must have 4 elements (NCHW)".into(),
));
}
if filter_shape.len() != 4 {
return Err(BackendError::InvalidArgument(
"filter_shape must have 4 elements (KCFHFW)".into(),
));
}
if output_shape.len() != 4 {
return Err(BackendError::InvalidArgument(
"output_shape must have 4 elements (NKOhOw)".into(),
));
}
if stride.len() != 2 {
return Err(BackendError::InvalidArgument(
"stride must have 2 elements [sh, sw]".into(),
));
}
if padding.len() != 2 {
return Err(BackendError::InvalidArgument(
"padding must have 2 elements [ph, pw]".into(),
));
}
let batch = input_shape[0];
let c_in = input_shape[1];
let h_in = input_shape[2];
let w_in = input_shape[3];
let k_out = filter_shape[0];
let fh = filter_shape[2];
let fw = filter_shape[3];
let oh = output_shape[2];
let ow = output_shape[3];
let sh = stride[0];
let sw = stride[1];
let ph = padding[0];
let pw = padding[1];
let in_elems: usize = input_shape.iter().product();
let f_elems: usize = filter_shape.iter().product();
let o_elems: usize = output_shape.iter().product();
if let Some((wg_x, wg_y)) = conv2d_gpu_dispatch_grid(batch, k_out, oh, ow) {
if let Some(dims) = conv2d_u32_dims(
batch, c_in, h_in, w_in, k_out, fh, fw, oh, ow, sh, sw, ph, pw,
) {
return self.conv2d_forward_gpu(
input_ptr, filter_ptr, output_ptr, dims, in_elems, f_elems, o_elems, wg_x, wg_y,
);
}
}
let mem = self.memory()?;
let mut in_bytes = vec![0u8; in_elems * 4];
let mut f_bytes = vec![0u8; f_elems * 4];
mem.copy_from_device(&mut in_bytes, input_ptr)
.map_err(BackendError::from)?;
mem.copy_from_device(&mut f_bytes, filter_ptr)
.map_err(BackendError::from)?;
let in_f32 = bytes_to_f32_vec(&in_bytes);
let f_f32 = bytes_to_f32_vec(&f_bytes);
let out_f32 = conv2d_cpu_reference(
&in_f32, &f_f32, batch, c_in, h_in, w_in, k_out, fh, fw, oh, ow, sh, sw, ph, pw,
);
let out_bytes = f32_slice_to_bytes(&out_f32);
mem.copy_to_device(output_ptr, &out_bytes)
.map_err(BackendError::from)?;
Ok(())
}
fn attention(
&self,
q_ptr: u64,
k_ptr: u64,
v_ptr: u64,
o_ptr: u64,
batch: usize,
heads: usize,
seq_q: usize,
seq_kv: usize,
head_dim: usize,
scale: f64,
causal: bool,
) -> BackendResult<()> {
self.check_init()?;
if seq_q == 0 || seq_kv == 0 || head_dim == 0 {
return Err(BackendError::InvalidArgument(
"seq_q, seq_kv, and head_dim must all be > 0".into(),
));
}
if scale <= 0.0 || !scale.is_finite() {
return Err(BackendError::InvalidArgument(format!(
"scale must be a positive finite number, got {scale}"
)));
}
let batch_heads = batch * heads;
let q_elems = batch_heads * seq_q * head_dim;
let kv_elems = batch_heads * seq_kv * head_dim;
let o_elems = q_elems;
let scale_f32 = scale as f32;
if let (Some(wg), Some(bh_u32), Some(seq_q_u32), Some(seq_kv_u32), Some(head_dim_u32)) = (
attention_gpu_dispatch_grid(batch_heads, seq_q),
u32::try_from(batch_heads).ok(),
u32::try_from(seq_q).ok(),
u32::try_from(seq_kv).ok(),
u32::try_from(head_dim).ok(),
) {
return self.attention_gpu(
q_ptr,
k_ptr,
v_ptr,
o_ptr,
bh_u32,
seq_q_u32,
seq_kv_u32,
head_dim_u32,
scale_f32,
causal,
q_elems,
kv_elems,
o_elems,
wg,
);
}
let mem = self.memory()?;
let mut q_bytes = vec![0u8; q_elems * 4];
let mut k_bytes = vec![0u8; kv_elems * 4];
let mut v_bytes = vec![0u8; kv_elems * 4];
mem.copy_from_device(&mut q_bytes, q_ptr)
.map_err(BackendError::from)?;
mem.copy_from_device(&mut k_bytes, k_ptr)
.map_err(BackendError::from)?;
mem.copy_from_device(&mut v_bytes, v_ptr)
.map_err(BackendError::from)?;
let q_f32 = bytes_to_f32_vec(&q_bytes);
let k_f32 = bytes_to_f32_vec(&k_bytes);
let v_f32 = bytes_to_f32_vec(&v_bytes);
let o_f32 = attention_cpu_reference(
&q_f32,
&k_f32,
&v_f32,
batch_heads,
seq_q,
seq_kv,
head_dim,
scale_f32,
causal,
);
let o_bytes = f32_slice_to_bytes(&o_f32);
mem.copy_to_device(o_ptr, &o_bytes)
.map_err(BackendError::from)?;
Ok(())
}
fn reduce(
&self,
op: ReduceOp,
input_ptr: u64,
output_ptr: u64,
shape: &[usize],
axis: usize,
) -> BackendResult<()> {
self.check_init()?;
if shape.is_empty() {
return Err(BackendError::InvalidArgument(
"shape must not be empty".into(),
));
}
if axis >= shape.len() {
return Err(BackendError::InvalidArgument(format!(
"axis {axis} is out of bounds for shape of length {}",
shape.len()
)));
}
if shape.len() != 1 {
return self.reduce_nd(op, input_ptr, output_ptr, shape, axis);
}
let n_elements = shape[0];
if n_elements == 0 {
return Ok(());
}
let dev = self.device()?;
let mem = self.memory()?;
let op_str = map_reduce_op(op);
let limits = gpu_limits();
let (wg_grid, _) = planner::plan_dispatch_1d(&limits, n_elements as u64, 256)
.map_err(BackendError::InvalidArgument)?;
if wg_grid.y != 1 {
return Err(BackendError::InvalidArgument(format!(
"reduce: {n_elements} elements need {} workgroups, which exceeds the \
single-axis dispatch capacity of this 1-D reduction kernel",
wg_grid.x as u64 * wg_grid.y as u64
)));
}
let wg_count = wg_grid.x;
let pass1_cached = self.cached_pipeline(
&format!("reduce_pass1:{op_str}"),
"oxicuda-reduce-pass1",
|| shader::reduction_wgsl(op_str),
)?;
let partial_buf = dev.device.create_buffer(&wgpu::BufferDescriptor {
label: Some("oxicuda-reduce-partial"),
size: (wg_count as u64) * 4, usage: wgpu::BufferUsages::STORAGE
| wgpu::BufferUsages::COPY_SRC
| wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
});
let mut p1_params = [0u8; 4];
p1_params[0..4].copy_from_slice(&(n_elements as u32).to_le_bytes());
let p1_uniform = dev.device.create_buffer(&wgpu::BufferDescriptor {
label: Some("oxicuda-reduce-p1-params"),
size: 4,
usage: wgpu::BufferUsages::UNIFORM | wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
});
dev.queue.write_buffer(&p1_uniform, 0, &p1_params);
let bgl1 = &pass1_cached.bind_group_layout;
let bg1 = {
let buffers = mem
.lock_buffers()
.map_err(|e| BackendError::DeviceError(e.to_string()))?;
let in_info = buffers.get(&input_ptr).ok_or_else(|| {
BackendError::InvalidArgument(format!("unknown handle {input_ptr}"))
})?;
let need_in = (n_elements as u64) * 4;
if in_info.size < need_in {
return Err(BackendError::InvalidArgument(format!(
"reduce: input buffer holds {} bytes, need {need_in} for {n_elements} f32 elements",
in_info.size
)));
}
dev.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some("oxicuda-reduce-pass1"),
layout: bgl1,
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: in_info.buffer.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: partial_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: p1_uniform.as_entire_binding(),
},
],
})
};
let pass2_cached = self.cached_pipeline(
&format!("reduce_pass2:{op_str}"),
"oxicuda-reduce-pass2",
|| shader::reduction_final_wgsl(op_str),
)?;
let mut p2_params = [0u8; 4];
p2_params[0..4].copy_from_slice(&wg_count.to_le_bytes());
let p2_uniform = dev.device.create_buffer(&wgpu::BufferDescriptor {
label: Some("oxicuda-reduce-p2-params"),
size: 4,
usage: wgpu::BufferUsages::UNIFORM | wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
});
dev.queue.write_buffer(&p2_uniform, 0, &p2_params);
let bgl2 = &pass2_cached.bind_group_layout;
let bg2 = {
let buffers = mem
.lock_buffers()
.map_err(|e| BackendError::DeviceError(e.to_string()))?;
let out_info = buffers.get(&output_ptr).ok_or_else(|| {
BackendError::InvalidArgument(format!("unknown handle {output_ptr}"))
})?;
if out_info.size < 4 {
return Err(BackendError::InvalidArgument(format!(
"reduce: output buffer holds {} bytes, need 4 for the scalar result",
out_info.size
)));
}
dev.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some("oxicuda-reduce-pass2"),
layout: bgl2,
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: partial_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: out_info.buffer.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: p2_uniform.as_entire_binding(),
},
],
})
};
let mut encoder = dev
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("oxicuda-reduce"),
});
{
let mut pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: Some("oxicuda-reduce-pass1"),
timestamp_writes: None,
});
pass.set_pipeline(&pass1_cached.pipeline);
pass.set_bind_group(0, &bg1, &[]);
pass.dispatch_workgroups(wg_count, 1, 1);
}
{
let mut pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: Some("oxicuda-reduce-pass2"),
timestamp_writes: None,
});
pass.set_pipeline(&pass2_cached.pipeline);
pass.set_bind_group(0, &bg2, &[]);
pass.dispatch_workgroups(1, 1, 1);
}
dev.queue.submit(std::iter::once(encoder.finish()));
if op == ReduceOp::Mean && n_elements > 1 {
let mut buf = [0u8; 4];
mem.copy_from_device(&mut buf, output_ptr)
.map_err(BackendError::from)?;
let val = f32::from_le_bytes(buf);
let mean = val / (n_elements as f32);
mem.copy_to_device(output_ptr, &mean.to_le_bytes())
.map_err(BackendError::from)?;
}
Ok(())
}
fn unary(&self, op: UnaryOp, input_ptr: u64, output_ptr: u64, n: usize) -> BackendResult<()> {
self.check_init()?;
if n == 0 {
return Ok(());
}
if input_ptr == output_ptr {
return Err(BackendError::InvalidArgument(
"unary: input_ptr and output_ptr must not alias (wgpu rejects binding the \
same buffer as both `read` and `read_write` within one dispatch); allocate \
a separate output buffer"
.into(),
));
}
let dev = self.device()?;
let mem = self.memory()?;
let (wg_grid, _) = planner::plan_dispatch_1d(&gpu_limits(), n as u64, 256)
.map_err(BackendError::InvalidArgument)?;
if wg_grid.y != 1 {
return Err(BackendError::InvalidArgument(format!(
"unary: {n} elements exceed the single-axis dispatch capacity of this kernel"
)));
}
let op_str = map_unary_op(op);
let pipeline_key = format!("unary:{op_str}");
let cached = self.cached_pipeline(&pipeline_key, "oxicuda-unary", || {
shader::elementwise_wgsl(op_str)
})?;
let bind_group = self.cached_bind_group(
dev,
mem,
&cached.bind_group_layout,
&pipeline_key,
&[input_ptr, output_ptr],
&[],
&[],
"oxicuda-unary",
)?;
let mut encoder = dev
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("oxicuda-unary"),
});
{
let mut pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: Some("oxicuda-unary"),
timestamp_writes: None,
});
pass.set_pipeline(&cached.pipeline);
pass.set_bind_group(0, &bind_group, &[]);
pass.dispatch_workgroups(wg_grid.x, 1, 1);
}
dev.queue.submit(std::iter::once(encoder.finish()));
Ok(())
}
fn binary(
&self,
op: BinaryOp,
a_ptr: u64,
b_ptr: u64,
output_ptr: u64,
n: usize,
) -> BackendResult<()> {
self.check_init()?;
if n == 0 {
return Ok(());
}
if a_ptr == output_ptr || b_ptr == output_ptr {
return Err(BackendError::InvalidArgument(
"binary: a_ptr/b_ptr must not alias output_ptr (wgpu rejects binding the \
same buffer as both `read` and `read_write` within one dispatch); allocate \
a separate output buffer"
.into(),
));
}
let dev = self.device()?;
let mem = self.memory()?;
let (wg_grid, _) = planner::plan_dispatch_1d(&gpu_limits(), n as u64, 256)
.map_err(BackendError::InvalidArgument)?;
if wg_grid.y != 1 {
return Err(BackendError::InvalidArgument(format!(
"binary: {n} elements exceed the single-axis dispatch capacity of this kernel"
)));
}
let op_str = map_binary_op(op);
let pipeline_key = format!("binary:{op_str}");
let cached = self.cached_pipeline(&pipeline_key, "oxicuda-binary", || {
shader::binary_wgsl(op_str)
})?;
let bind_group = self.cached_bind_group(
dev,
mem,
&cached.bind_group_layout,
&pipeline_key,
&[a_ptr, b_ptr, output_ptr],
&[],
&[],
"oxicuda-binary",
)?;
let mut encoder = dev
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("oxicuda-binary"),
});
{
let mut pass = encoder.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: Some("oxicuda-binary"),
timestamp_writes: None,
});
pass.set_pipeline(&cached.pipeline);
pass.set_bind_group(0, &bind_group, &[]);
pass.dispatch_workgroups(wg_grid.x, 1, 1);
}
dev.queue.submit(std::iter::once(encoder.finish()));
Ok(())
}
fn synchronize(&self) -> BackendResult<()> {
self.check_init()?;
if let Some(dev) = &self.device {
crate::memory::poll_result_to_webgpu_result(
dev.device.poll(wgpu::PollType::wait_indefinitely()),
)
.map_err(BackendError::from)?;
if let Some(msg) = dev.poll_error() {
return Err(BackendError::from(
crate::error::WebGpuError::UncapturedError(msg),
));
}
}
Ok(())
}
fn alloc(&self, bytes: usize) -> BackendResult<u64> {
self.check_init()?;
if bytes == 0 {
return Err(BackendError::InvalidArgument(
"cannot allocate 0 bytes".into(),
));
}
self.memory()?.alloc(bytes).map_err(BackendError::from)
}
fn free(&self, ptr: u64) -> BackendResult<()> {
self.check_init()?;
self.evict_bind_group_cache(ptr)?;
self.memory()?.free(ptr).map_err(BackendError::from)
}
fn copy_htod(&self, dst: u64, src: &[u8]) -> BackendResult<()> {
self.check_init()?;
if src.is_empty() {
return Ok(());
}
self.memory()?
.copy_to_device(dst, src)
.map_err(BackendError::from)
}
fn copy_dtoh(&self, dst: &mut [u8], src: u64) -> BackendResult<()> {
self.check_init()?;
if dst.is_empty() {
return Ok(());
}
self.memory()?
.copy_from_device(dst, src)
.map_err(BackendError::from)
}
}
fn bytes_to_f32_vec(bytes: &[u8]) -> Vec<f32> {
bytes
.chunks_exact(4)
.map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
.collect()
}
fn f32_slice_to_bytes(data: &[f32]) -> Vec<u8> {
data.iter().flat_map(|v| v.to_le_bytes()).collect()
}
#[cfg(test)]
#[path = "backend_tests.rs"]
mod tests;