941 lines
29 KiB
Rust
941 lines
29 KiB
Rust
//! WebGPU Backend for Browser GPU Acceleration
|
|
//!
|
|
//! This module provides GPU-accelerated inference using the WebGPU API,
|
|
//! enabling 10-100x speedup compared to CPU-only WASM execution.
|
|
//!
|
|
//! ## Features
|
|
//!
|
|
//! - **GPU Matrix Operations**: GEMM, attention, layer norm on GPU
|
|
//! - **Shader Compilation**: Dynamic WGSL compute shaders
|
|
//! - **Memory Management**: Efficient GPU buffer pooling
|
|
//! - **Fallback**: Automatic CPU fallback when WebGPU unavailable
|
|
//!
|
|
//! ## Browser Support
|
|
//!
|
|
//! - Chrome 113+ (full support)
|
|
//! - Firefox 121+ (behind flag)
|
|
//! - Safari 17+ (partial support)
|
|
//! - Edge 113+ (full support)
|
|
|
|
use js_sys::Float32Array;
|
|
use serde::{Deserialize, Serialize};
|
|
use std::collections::HashMap;
|
|
use std::sync::atomic::{AtomicU64, Ordering};
|
|
use wasm_bindgen::prelude::*;
|
|
use wasm_bindgen_futures::JsFuture;
|
|
|
|
/// WebGPU availability status
|
|
#[wasm_bindgen]
|
|
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
|
pub enum WebGPUStatus {
|
|
/// WebGPU is available and initialized
|
|
Available,
|
|
/// WebGPU is not supported in this browser
|
|
NotSupported,
|
|
/// WebGPU adapter request failed
|
|
AdapterFailed,
|
|
/// WebGPU device request failed
|
|
DeviceFailed,
|
|
/// WebGPU is available but not yet initialized
|
|
NotInitialized,
|
|
}
|
|
|
|
/// WebGPU adapter information
|
|
#[wasm_bindgen(getter_with_clone)]
|
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
pub struct WebGPUAdapterInfo {
|
|
/// Adapter vendor
|
|
pub vendor: String,
|
|
/// Adapter architecture
|
|
pub architecture: String,
|
|
/// Adapter device name
|
|
pub device: String,
|
|
/// Adapter description
|
|
pub description: String,
|
|
/// Is fallback (software) adapter
|
|
pub is_fallback: bool,
|
|
}
|
|
|
|
#[wasm_bindgen]
|
|
impl WebGPUAdapterInfo {
|
|
/// Returns the adapter vendor name.
|
|
#[wasm_bindgen(getter)]
|
|
pub fn vendor(&self) -> String {
|
|
self.vendor.clone()
|
|
}
|
|
|
|
/// Returns the adapter architecture.
|
|
#[wasm_bindgen(getter)]
|
|
pub fn architecture(&self) -> String {
|
|
self.architecture.clone()
|
|
}
|
|
|
|
/// Returns the adapter device name.
|
|
#[wasm_bindgen(getter)]
|
|
pub fn device(&self) -> String {
|
|
self.device.clone()
|
|
}
|
|
|
|
/// Returns the adapter description.
|
|
#[wasm_bindgen(getter)]
|
|
pub fn description(&self) -> String {
|
|
self.description.clone()
|
|
}
|
|
|
|
/// Returns whether this is a fallback (software) adapter.
|
|
#[wasm_bindgen(getter)]
|
|
pub fn is_fallback(&self) -> bool {
|
|
self.is_fallback
|
|
}
|
|
}
|
|
|
|
/// WebGPU device limits
|
|
#[wasm_bindgen]
|
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
|
pub struct WebGPULimits {
|
|
/// Maximum buffer size
|
|
pub max_buffer_size: u64,
|
|
/// Maximum storage buffer binding size
|
|
pub max_storage_buffer_binding_size: u64,
|
|
/// Maximum compute workgroup size X
|
|
pub max_compute_workgroup_size_x: u32,
|
|
/// Maximum compute workgroup size Y
|
|
pub max_compute_workgroup_size_y: u32,
|
|
/// Maximum compute workgroup size Z
|
|
pub max_compute_workgroup_size_z: u32,
|
|
/// Maximum compute invocations per workgroup
|
|
pub max_compute_invocations_per_workgroup: u32,
|
|
/// Maximum bind groups
|
|
pub max_bind_groups: u32,
|
|
}
|
|
|
|
impl Default for WebGPULimits {
|
|
fn default() -> Self {
|
|
Self {
|
|
max_buffer_size: 256 * 1024 * 1024, // 256 MB
|
|
max_storage_buffer_binding_size: 128 * 1024 * 1024, // 128 MB
|
|
max_compute_workgroup_size_x: 256,
|
|
max_compute_workgroup_size_y: 256,
|
|
max_compute_workgroup_size_z: 64,
|
|
max_compute_invocations_per_workgroup: 256,
|
|
max_bind_groups: 4,
|
|
}
|
|
}
|
|
}
|
|
|
|
/// WebGPU context for inference
|
|
#[wasm_bindgen]
|
|
pub struct WebGPUContext {
|
|
status: WebGPUStatus,
|
|
adapter_info: Option<WebGPUAdapterInfo>,
|
|
limits: WebGPULimits,
|
|
// Internal handles (stored as JsValue since we can't store web_sys types directly)
|
|
device_handle: Option<JsValue>,
|
|
queue_handle: Option<JsValue>,
|
|
// Shader cache
|
|
shader_cache: HashMap<String, JsValue>,
|
|
// Buffer pool
|
|
buffer_pool: Vec<WebGPUBuffer>,
|
|
// Stats
|
|
total_compute_time_ms: f64,
|
|
total_memory_bytes: u64,
|
|
}
|
|
|
|
/// GPU buffer wrapper
|
|
#[derive(Debug, Clone)]
|
|
struct WebGPUBuffer {
|
|
id: u64,
|
|
size: u64,
|
|
usage: WebGPUBufferUsage,
|
|
handle: JsValue,
|
|
in_use: bool,
|
|
}
|
|
|
|
/// Buffer usage flags
|
|
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
|
pub enum WebGPUBufferUsage {
|
|
/// Read-only storage buffer
|
|
Storage,
|
|
/// Read-write storage buffer
|
|
StorageReadWrite,
|
|
/// Uniform buffer
|
|
Uniform,
|
|
/// Staging buffer for CPU readback
|
|
Staging,
|
|
}
|
|
|
|
static BUFFER_ID_COUNTER: AtomicU64 = AtomicU64::new(0);
|
|
|
|
#[wasm_bindgen]
|
|
impl WebGPUContext {
|
|
/// Check if WebGPU is supported in the current browser
|
|
#[wasm_bindgen]
|
|
pub fn is_supported() -> bool {
|
|
let global = js_sys::global();
|
|
js_sys::Reflect::get(&global, &"navigator".into())
|
|
.ok()
|
|
.and_then(|nav| js_sys::Reflect::get(&nav, &"gpu".into()).ok())
|
|
.is_some_and(|gpu| !gpu.is_undefined())
|
|
}
|
|
|
|
/// Create and initialize WebGPU context
|
|
#[wasm_bindgen(constructor)]
|
|
pub async fn new() -> Result<WebGPUContext, JsError> {
|
|
if !Self::is_supported() {
|
|
return Ok(WebGPUContext {
|
|
status: WebGPUStatus::NotSupported,
|
|
adapter_info: None,
|
|
limits: WebGPULimits::default(),
|
|
device_handle: None,
|
|
queue_handle: None,
|
|
shader_cache: HashMap::new(),
|
|
buffer_pool: Vec::new(),
|
|
total_compute_time_ms: 0.0,
|
|
total_memory_bytes: 0,
|
|
});
|
|
}
|
|
|
|
// Get navigator.gpu
|
|
let global = js_sys::global();
|
|
let navigator = js_sys::Reflect::get(&global, &"navigator".into())
|
|
.map_err(|e| JsError::new(&format!("Failed to get navigator: {:?}", e)))?;
|
|
let gpu = js_sys::Reflect::get(&navigator, &"gpu".into())
|
|
.map_err(|e| JsError::new(&format!("Failed to get GPU: {:?}", e)))?;
|
|
|
|
// Request adapter
|
|
let adapter_promise = js_sys::Reflect::get(&gpu, &"requestAdapter".into())
|
|
.map_err(|e| JsError::new(&format!("Failed to get requestAdapter: {:?}", e)))?;
|
|
|
|
let request_adapter_fn = adapter_promise
|
|
.dyn_ref::<js_sys::Function>()
|
|
.ok_or_else(|| JsError::new("requestAdapter is not a function"))?;
|
|
|
|
// Request with high-performance preference
|
|
let options = js_sys::Object::new();
|
|
js_sys::Reflect::set(
|
|
&options,
|
|
&"powerPreference".into(),
|
|
&"high-performance".into(),
|
|
)
|
|
.map_err(|e| JsError::new(&format!("Failed to set power preference: {:?}", e)))?;
|
|
|
|
let adapter_promise = request_adapter_fn
|
|
.call1(&gpu, &options)
|
|
.map_err(|e| JsError::new(&format!("Failed to call requestAdapter: {:?}", e)))?;
|
|
|
|
let adapter = JsFuture::from(js_sys::Promise::from(adapter_promise))
|
|
.await
|
|
.map_err(|e| JsError::new(&format!("Adapter request failed: {:?}", e)))?;
|
|
|
|
if adapter.is_null() || adapter.is_undefined() {
|
|
return Ok(WebGPUContext {
|
|
status: WebGPUStatus::AdapterFailed,
|
|
adapter_info: None,
|
|
limits: WebGPULimits::default(),
|
|
device_handle: None,
|
|
queue_handle: None,
|
|
shader_cache: HashMap::new(),
|
|
buffer_pool: Vec::new(),
|
|
total_compute_time_ms: 0.0,
|
|
total_memory_bytes: 0,
|
|
});
|
|
}
|
|
|
|
// Get adapter info
|
|
let adapter_info = Self::get_adapter_info(&adapter)?;
|
|
|
|
// Request device
|
|
let request_device_fn = js_sys::Reflect::get(&adapter, &"requestDevice".into())
|
|
.map_err(|e| JsError::new(&format!("Failed to get requestDevice: {:?}", e)))?
|
|
.dyn_ref::<js_sys::Function>()
|
|
.ok_or_else(|| JsError::new("requestDevice is not a function"))?
|
|
.clone();
|
|
|
|
let device_promise = request_device_fn
|
|
.call0(&adapter)
|
|
.map_err(|e| JsError::new(&format!("Failed to call requestDevice: {:?}", e)))?;
|
|
|
|
let device = JsFuture::from(js_sys::Promise::from(device_promise))
|
|
.await
|
|
.map_err(|e| JsError::new(&format!("Device request failed: {:?}", e)))?;
|
|
|
|
if device.is_null() || device.is_undefined() {
|
|
return Ok(WebGPUContext {
|
|
status: WebGPUStatus::DeviceFailed,
|
|
adapter_info: Some(adapter_info),
|
|
limits: WebGPULimits::default(),
|
|
device_handle: None,
|
|
queue_handle: None,
|
|
shader_cache: HashMap::new(),
|
|
buffer_pool: Vec::new(),
|
|
total_compute_time_ms: 0.0,
|
|
total_memory_bytes: 0,
|
|
});
|
|
}
|
|
|
|
// Get device queue
|
|
let queue = js_sys::Reflect::get(&device, &"queue".into())
|
|
.map_err(|e| JsError::new(&format!("Failed to get queue: {:?}", e)))?;
|
|
|
|
// Get device limits
|
|
let limits = Self::get_device_limits(&device)?;
|
|
|
|
Ok(WebGPUContext {
|
|
status: WebGPUStatus::Available,
|
|
adapter_info: Some(adapter_info),
|
|
limits,
|
|
device_handle: Some(device),
|
|
queue_handle: Some(queue),
|
|
shader_cache: HashMap::new(),
|
|
buffer_pool: Vec::new(),
|
|
total_compute_time_ms: 0.0,
|
|
total_memory_bytes: 0,
|
|
})
|
|
}
|
|
|
|
fn get_adapter_info(adapter: &JsValue) -> Result<WebGPUAdapterInfo, JsError> {
|
|
// Try to get adapter info (async in WebGPU spec, but we'll use sync access if available)
|
|
let info = js_sys::Reflect::get(adapter, &"info".into()).ok();
|
|
|
|
if let Some(info) = info {
|
|
if !info.is_undefined() {
|
|
let vendor = js_sys::Reflect::get(&info, &"vendor".into())
|
|
.ok()
|
|
.and_then(|v| v.as_string())
|
|
.unwrap_or_else(|| "unknown".to_string());
|
|
let architecture = js_sys::Reflect::get(&info, &"architecture".into())
|
|
.ok()
|
|
.and_then(|v| v.as_string())
|
|
.unwrap_or_else(|| "unknown".to_string());
|
|
let device = js_sys::Reflect::get(&info, &"device".into())
|
|
.ok()
|
|
.and_then(|v| v.as_string())
|
|
.unwrap_or_else(|| "unknown".to_string());
|
|
let description = js_sys::Reflect::get(&info, &"description".into())
|
|
.ok()
|
|
.and_then(|v| v.as_string())
|
|
.unwrap_or_else(|| "unknown".to_string());
|
|
|
|
return Ok(WebGPUAdapterInfo {
|
|
vendor,
|
|
architecture,
|
|
device,
|
|
description,
|
|
is_fallback: false,
|
|
});
|
|
}
|
|
}
|
|
|
|
// Fallback: check if it's a fallback adapter
|
|
let is_fallback = js_sys::Reflect::get(adapter, &"isFallbackAdapter".into())
|
|
.ok()
|
|
.and_then(|v| v.as_bool())
|
|
.unwrap_or(false);
|
|
|
|
Ok(WebGPUAdapterInfo {
|
|
vendor: "unknown".to_string(),
|
|
architecture: "unknown".to_string(),
|
|
device: "unknown".to_string(),
|
|
description: "WebGPU Adapter".to_string(),
|
|
is_fallback,
|
|
})
|
|
}
|
|
|
|
fn get_device_limits(device: &JsValue) -> Result<WebGPULimits, JsError> {
|
|
let limits = js_sys::Reflect::get(device, &"limits".into())
|
|
.map_err(|e| JsError::new(&format!("Failed to get limits: {:?}", e)))?;
|
|
|
|
if limits.is_undefined() {
|
|
return Ok(WebGPULimits::default());
|
|
}
|
|
|
|
let get_limit = |name: &str, default: u64| -> u64 {
|
|
js_sys::Reflect::get(&limits, &name.into())
|
|
.ok()
|
|
.and_then(|v| v.as_f64())
|
|
.map_or(default, |v| v as u64)
|
|
};
|
|
|
|
Ok(WebGPULimits {
|
|
max_buffer_size: get_limit("maxBufferSize", 256 * 1024 * 1024),
|
|
max_storage_buffer_binding_size: get_limit(
|
|
"maxStorageBufferBindingSize",
|
|
128 * 1024 * 1024,
|
|
),
|
|
max_compute_workgroup_size_x: get_limit("maxComputeWorkgroupSizeX", 256) as u32,
|
|
max_compute_workgroup_size_y: get_limit("maxComputeWorkgroupSizeY", 256) as u32,
|
|
max_compute_workgroup_size_z: get_limit("maxComputeWorkgroupSizeZ", 64) as u32,
|
|
max_compute_invocations_per_workgroup: get_limit(
|
|
"maxComputeInvocationsPerWorkgroup",
|
|
256,
|
|
) as u32,
|
|
max_bind_groups: get_limit("maxBindGroups", 4) as u32,
|
|
})
|
|
}
|
|
|
|
/// Get WebGPU status
|
|
#[wasm_bindgen(getter)]
|
|
pub fn status(&self) -> WebGPUStatus {
|
|
self.status
|
|
}
|
|
|
|
/// Check if WebGPU is available and ready
|
|
#[wasm_bindgen]
|
|
pub fn is_available(&self) -> bool {
|
|
self.status == WebGPUStatus::Available
|
|
}
|
|
|
|
/// Get adapter info
|
|
#[wasm_bindgen]
|
|
pub fn adapter_info(&self) -> Option<WebGPUAdapterInfo> {
|
|
self.adapter_info.clone()
|
|
}
|
|
|
|
/// Get device limits
|
|
#[wasm_bindgen]
|
|
pub fn limits(&self) -> WebGPULimits {
|
|
self.limits.clone()
|
|
}
|
|
|
|
/// Get total compute time in milliseconds
|
|
#[wasm_bindgen(getter)]
|
|
pub fn total_compute_time_ms(&self) -> f64 {
|
|
self.total_compute_time_ms
|
|
}
|
|
|
|
/// Get total GPU memory usage
|
|
#[wasm_bindgen(getter)]
|
|
pub fn total_memory_bytes(&self) -> u64 {
|
|
self.total_memory_bytes
|
|
}
|
|
|
|
/// Create a GPU buffer with data
|
|
#[wasm_bindgen]
|
|
pub fn create_buffer_with_data(
|
|
&mut self,
|
|
data: &[f32],
|
|
read_only: bool,
|
|
) -> Result<u64, JsError> {
|
|
let device = self
|
|
.device_handle
|
|
.as_ref()
|
|
.ok_or_else(|| JsError::new("WebGPU device not initialized"))?;
|
|
|
|
let size = (data.len() * 4) as u64;
|
|
|
|
// Helper to convert JsValue error to JsError
|
|
fn js_err(e: JsValue) -> JsError {
|
|
JsError::new(&format!("{:?}", e))
|
|
}
|
|
|
|
// Create buffer descriptor
|
|
let buffer_desc = js_sys::Object::new();
|
|
js_sys::Reflect::set(&buffer_desc, &"size".into(), &JsValue::from(size)).map_err(js_err)?;
|
|
|
|
// GPUBufferUsage.STORAGE = 0x80, GPUBufferUsage.COPY_DST = 0x08, GPUBufferUsage.COPY_SRC = 0x04
|
|
let usage = if read_only {
|
|
0x80 | 0x08
|
|
} else {
|
|
0x80 | 0x08 | 0x04
|
|
};
|
|
js_sys::Reflect::set(&buffer_desc, &"usage".into(), &JsValue::from(usage))
|
|
.map_err(js_err)?;
|
|
js_sys::Reflect::set(&buffer_desc, &"mappedAtCreation".into(), &JsValue::TRUE)
|
|
.map_err(js_err)?;
|
|
|
|
// Create buffer
|
|
let create_buffer_fn = js_sys::Reflect::get(device, &"createBuffer".into())
|
|
.map_err(js_err)?
|
|
.dyn_ref::<js_sys::Function>()
|
|
.ok_or_else(|| JsError::new("createBuffer is not a function"))?
|
|
.clone();
|
|
|
|
let buffer = create_buffer_fn
|
|
.call1(device, &buffer_desc)
|
|
.map_err(js_err)?;
|
|
|
|
// Write data
|
|
let get_mapped_range_fn = js_sys::Reflect::get(&buffer, &"getMappedRange".into())
|
|
.map_err(js_err)?
|
|
.dyn_ref::<js_sys::Function>()
|
|
.ok_or_else(|| JsError::new("getMappedRange is not a function"))?
|
|
.clone();
|
|
|
|
let array_buffer = get_mapped_range_fn.call0(&buffer).map_err(js_err)?;
|
|
let typed_array = Float32Array::new(&array_buffer);
|
|
typed_array.copy_from(data);
|
|
|
|
// Unmap buffer
|
|
let unmap_fn = js_sys::Reflect::get(&buffer, &"unmap".into())
|
|
.map_err(js_err)?
|
|
.dyn_ref::<js_sys::Function>()
|
|
.ok_or_else(|| JsError::new("unmap is not a function"))?
|
|
.clone();
|
|
unmap_fn.call0(&buffer).map_err(js_err)?;
|
|
|
|
// Track buffer
|
|
let buffer_id = BUFFER_ID_COUNTER.fetch_add(1, Ordering::SeqCst);
|
|
self.buffer_pool.push(WebGPUBuffer {
|
|
id: buffer_id,
|
|
size,
|
|
usage: if read_only {
|
|
WebGPUBufferUsage::Storage
|
|
} else {
|
|
WebGPUBufferUsage::StorageReadWrite
|
|
},
|
|
handle: buffer,
|
|
in_use: true,
|
|
});
|
|
self.total_memory_bytes += size;
|
|
|
|
Ok(buffer_id)
|
|
}
|
|
|
|
/// Create an empty GPU buffer
|
|
#[wasm_bindgen]
|
|
pub fn create_buffer(&mut self, size: u64, read_only: bool) -> Result<u64, JsError> {
|
|
let device = self
|
|
.device_handle
|
|
.as_ref()
|
|
.ok_or_else(|| JsError::new("WebGPU device not initialized"))?;
|
|
|
|
// Helper to convert JsValue error to JsError
|
|
fn js_err(e: JsValue) -> JsError {
|
|
JsError::new(&format!("{:?}", e))
|
|
}
|
|
|
|
let buffer_desc = js_sys::Object::new();
|
|
js_sys::Reflect::set(&buffer_desc, &"size".into(), &JsValue::from(size)).map_err(js_err)?;
|
|
|
|
let usage = if read_only {
|
|
0x80 | 0x08
|
|
} else {
|
|
0x80 | 0x08 | 0x04
|
|
};
|
|
js_sys::Reflect::set(&buffer_desc, &"usage".into(), &JsValue::from(usage))
|
|
.map_err(js_err)?;
|
|
|
|
let create_buffer_fn = js_sys::Reflect::get(device, &"createBuffer".into())
|
|
.map_err(js_err)?
|
|
.dyn_ref::<js_sys::Function>()
|
|
.ok_or_else(|| JsError::new("createBuffer is not a function"))?
|
|
.clone();
|
|
|
|
let buffer = create_buffer_fn
|
|
.call1(device, &buffer_desc)
|
|
.map_err(js_err)?;
|
|
|
|
let buffer_id = BUFFER_ID_COUNTER.fetch_add(1, Ordering::SeqCst);
|
|
self.buffer_pool.push(WebGPUBuffer {
|
|
id: buffer_id,
|
|
size,
|
|
usage: if read_only {
|
|
WebGPUBufferUsage::Storage
|
|
} else {
|
|
WebGPUBufferUsage::StorageReadWrite
|
|
},
|
|
handle: buffer,
|
|
in_use: true,
|
|
});
|
|
self.total_memory_bytes += size;
|
|
|
|
Ok(buffer_id)
|
|
}
|
|
|
|
/// Release a GPU buffer
|
|
#[wasm_bindgen]
|
|
pub fn release_buffer(&mut self, buffer_id: u64) -> Result<(), JsError> {
|
|
if let Some(idx) = self.buffer_pool.iter().position(|b| b.id == buffer_id) {
|
|
let buffer = &self.buffer_pool[idx];
|
|
|
|
// Destroy buffer
|
|
if let Ok(destroy_fn) = js_sys::Reflect::get(&buffer.handle, &"destroy".into()) {
|
|
if let Some(destroy_fn) = destroy_fn.dyn_ref::<js_sys::Function>() {
|
|
let _ = destroy_fn.call0(&buffer.handle);
|
|
}
|
|
}
|
|
|
|
self.total_memory_bytes = self.total_memory_bytes.saturating_sub(buffer.size);
|
|
self.buffer_pool.remove(idx);
|
|
}
|
|
Ok(())
|
|
}
|
|
|
|
/// Get number of active buffers
|
|
#[wasm_bindgen]
|
|
pub fn active_buffer_count(&self) -> usize {
|
|
self.buffer_pool.len()
|
|
}
|
|
}
|
|
|
|
//==============================================================================
|
|
// WGSL Compute Shaders for Inference
|
|
//==============================================================================
|
|
|
|
/// Matrix multiplication shader (GEMM)
|
|
pub const GEMM_SHADER: &str = r"
|
|
@group(0) @binding(0) var<storage, read> a: array<f32>;
|
|
@group(0) @binding(1) var<storage, read> b: array<f32>;
|
|
@group(0) @binding(2) var<storage, read_write> c: array<f32>;
|
|
|
|
struct Dimensions {
|
|
M: u32,
|
|
N: u32,
|
|
K: u32,
|
|
}
|
|
@group(0) @binding(3) var<uniform> dims: Dimensions;
|
|
|
|
const TILE_SIZE: u32 = 16u;
|
|
|
|
@compute @workgroup_size(16, 16)
|
|
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
|
|
let row = global_id.x;
|
|
let col = global_id.y;
|
|
|
|
if (row >= dims.M || col >= dims.N) {
|
|
return;
|
|
}
|
|
|
|
var sum: f32 = 0.0;
|
|
for (var k: u32 = 0u; k < dims.K; k = k + 1u) {
|
|
let a_idx = row * dims.K + k;
|
|
let b_idx = k * dims.N + col;
|
|
sum = sum + a[a_idx] * b[b_idx];
|
|
}
|
|
|
|
let c_idx = row * dims.N + col;
|
|
c[c_idx] = sum;
|
|
}
|
|
";
|
|
|
|
/// Softmax shader
|
|
pub const SOFTMAX_SHADER: &str = r"
|
|
@group(0) @binding(0) var<storage, read> input: array<f32>;
|
|
@group(0) @binding(1) var<storage, read_write> output: array<f32>;
|
|
|
|
struct SoftmaxParams {
|
|
rows: u32,
|
|
cols: u32,
|
|
}
|
|
@group(0) @binding(2) var<uniform> params: SoftmaxParams;
|
|
|
|
@compute @workgroup_size(256)
|
|
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
|
|
let row = global_id.x;
|
|
|
|
if (row >= params.rows) {
|
|
return;
|
|
}
|
|
|
|
let row_start = row * params.cols;
|
|
|
|
// Find max for numerical stability
|
|
var max_val: f32 = input[row_start];
|
|
for (var i: u32 = 1u; i < params.cols; i = i + 1u) {
|
|
max_val = max(max_val, input[row_start + i]);
|
|
}
|
|
|
|
// Compute exp and sum
|
|
var sum: f32 = 0.0;
|
|
for (var i: u32 = 0u; i < params.cols; i = i + 1u) {
|
|
let exp_val = exp(input[row_start + i] - max_val);
|
|
output[row_start + i] = exp_val;
|
|
sum = sum + exp_val;
|
|
}
|
|
|
|
// Normalize
|
|
for (var i: u32 = 0u; i < params.cols; i = i + 1u) {
|
|
output[row_start + i] = output[row_start + i] / sum;
|
|
}
|
|
}
|
|
";
|
|
|
|
/// Layer normalization shader
|
|
pub const LAYER_NORM_SHADER: &str = r"
|
|
@group(0) @binding(0) var<storage, read> input: array<f32>;
|
|
@group(0) @binding(1) var<storage, read> gamma: array<f32>;
|
|
@group(0) @binding(2) var<storage, read> beta: array<f32>;
|
|
@group(0) @binding(3) var<storage, read_write> output: array<f32>;
|
|
|
|
struct LayerNormParams {
|
|
batch_size: u32,
|
|
hidden_size: u32,
|
|
eps: f32,
|
|
}
|
|
@group(0) @binding(4) var<uniform> params: LayerNormParams;
|
|
|
|
@compute @workgroup_size(256)
|
|
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
|
|
let batch_idx = global_id.x;
|
|
|
|
if (batch_idx >= params.batch_size) {
|
|
return;
|
|
}
|
|
|
|
let offset = batch_idx * params.hidden_size;
|
|
|
|
// Compute mean
|
|
var sum: f32 = 0.0;
|
|
for (var i: u32 = 0u; i < params.hidden_size; i = i + 1u) {
|
|
sum = sum + input[offset + i];
|
|
}
|
|
let mean = sum / f32(params.hidden_size);
|
|
|
|
// Compute variance
|
|
var var_sum: f32 = 0.0;
|
|
for (var i: u32 = 0u; i < params.hidden_size; i = i + 1u) {
|
|
let diff = input[offset + i] - mean;
|
|
var_sum = var_sum + diff * diff;
|
|
}
|
|
let variance = var_sum / f32(params.hidden_size);
|
|
let inv_std = 1.0 / sqrt(variance + params.eps);
|
|
|
|
// Normalize and apply affine transform
|
|
for (var i: u32 = 0u; i < params.hidden_size; i = i + 1u) {
|
|
let normalized = (input[offset + i] - mean) * inv_std;
|
|
output[offset + i] = normalized * gamma[i] + beta[i];
|
|
}
|
|
}
|
|
";
|
|
|
|
/// GELU activation shader
|
|
pub const GELU_SHADER: &str = r"
|
|
@group(0) @binding(0) var<storage, read> input: array<f32>;
|
|
@group(0) @binding(1) var<storage, read_write> output: array<f32>;
|
|
|
|
struct GELUParams {
|
|
size: u32,
|
|
}
|
|
@group(0) @binding(2) var<uniform> params: GELUParams;
|
|
|
|
const SQRT_2_OVER_PI: f32 = 0.7978845608;
|
|
const GELU_CONST: f32 = 0.044715;
|
|
|
|
@compute @workgroup_size(256)
|
|
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
|
|
let idx = global_id.x;
|
|
|
|
if (idx >= params.size) {
|
|
return;
|
|
}
|
|
|
|
let x = input[idx];
|
|
// GELU(x) = 0.5 * x * (1 + tanh(sqrt(2/pi) * (x + 0.044715 * x^3)))
|
|
let inner = SQRT_2_OVER_PI * (x + GELU_CONST * x * x * x);
|
|
output[idx] = 0.5 * x * (1.0 + tanh(inner));
|
|
}
|
|
";
|
|
|
|
/// Scaled dot-product attention shader
|
|
pub const ATTENTION_SHADER: &str = r"
|
|
@group(0) @binding(0) var<storage, read> query: array<f32>;
|
|
@group(0) @binding(1) var<storage, read> key: array<f32>;
|
|
@group(0) @binding(2) var<storage, read> value: array<f32>;
|
|
@group(0) @binding(3) var<storage, read_write> output: array<f32>;
|
|
|
|
struct AttentionParams {
|
|
batch_size: u32,
|
|
num_heads: u32,
|
|
seq_len: u32,
|
|
head_dim: u32,
|
|
scale: f32,
|
|
causal: u32,
|
|
}
|
|
@group(0) @binding(4) var<uniform> params: AttentionParams;
|
|
|
|
// Shared memory for attention scores
|
|
var<workgroup> scores: array<f32, 1024>; // Max seq_len * seq_len for one head
|
|
|
|
@compute @workgroup_size(16, 16)
|
|
fn main(@builtin(global_invocation_id) global_id: vec3<u32>,
|
|
@builtin(local_invocation_id) local_id: vec3<u32>,
|
|
@builtin(workgroup_id) workgroup_id: vec3<u32>) {
|
|
|
|
let batch_idx = workgroup_id.z;
|
|
let head_idx = workgroup_id.y;
|
|
let q_pos = global_id.x;
|
|
let k_pos = global_id.y;
|
|
|
|
if (batch_idx >= params.batch_size || head_idx >= params.num_heads ||
|
|
q_pos >= params.seq_len || k_pos >= params.seq_len) {
|
|
return;
|
|
}
|
|
|
|
// Apply causal mask
|
|
if (params.causal > 0u && k_pos > q_pos) {
|
|
scores[q_pos * params.seq_len + k_pos] = -1e9;
|
|
return;
|
|
}
|
|
|
|
// Compute Q @ K^T for this position
|
|
let head_offset = (batch_idx * params.num_heads + head_idx) * params.seq_len * params.head_dim;
|
|
let q_offset = head_offset + q_pos * params.head_dim;
|
|
let k_offset = head_offset + k_pos * params.head_dim;
|
|
|
|
var dot: f32 = 0.0;
|
|
for (var d: u32 = 0u; d < params.head_dim; d = d + 1u) {
|
|
dot = dot + query[q_offset + d] * key[k_offset + d];
|
|
}
|
|
|
|
scores[q_pos * params.seq_len + k_pos] = dot * params.scale;
|
|
}
|
|
";
|
|
|
|
/// RoPE (Rotary Position Embedding) shader
|
|
pub const ROPE_SHADER: &str = r"
|
|
@group(0) @binding(0) var<storage, read_write> x: array<f32>;
|
|
|
|
struct RoPEParams {
|
|
batch_size: u32,
|
|
seq_len: u32,
|
|
num_heads: u32,
|
|
head_dim: u32,
|
|
theta_base: f32,
|
|
}
|
|
@group(0) @binding(1) var<uniform> params: RoPEParams;
|
|
|
|
@compute @workgroup_size(64)
|
|
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
|
|
let batch_idx = global_id.z;
|
|
let pos = global_id.y;
|
|
let head_idx = global_id.x / (params.head_dim / 2u);
|
|
let dim_pair = global_id.x % (params.head_dim / 2u);
|
|
|
|
if (batch_idx >= params.batch_size || pos >= params.seq_len ||
|
|
head_idx >= params.num_heads || dim_pair >= params.head_dim / 2u) {
|
|
return;
|
|
}
|
|
|
|
// Compute rotation angle
|
|
let freq = pow(params.theta_base, -f32(2u * dim_pair) / f32(params.head_dim));
|
|
let angle = f32(pos) * freq;
|
|
let cos_angle = cos(angle);
|
|
let sin_angle = sin(angle);
|
|
|
|
// Get indices for the pair
|
|
let base_offset = (batch_idx * params.num_heads * params.seq_len * params.head_dim) +
|
|
(head_idx * params.seq_len * params.head_dim) +
|
|
(pos * params.head_dim);
|
|
let idx1 = base_offset + 2u * dim_pair;
|
|
let idx2 = base_offset + 2u * dim_pair + 1u;
|
|
|
|
// Apply rotation
|
|
let x1 = x[idx1];
|
|
let x2 = x[idx2];
|
|
x[idx1] = x1 * cos_angle - x2 * sin_angle;
|
|
x[idx2] = x1 * sin_angle + x2 * cos_angle;
|
|
}
|
|
";
|
|
|
|
/// Elementwise add shader
|
|
pub const ADD_SHADER: &str = r"
|
|
@group(0) @binding(0) var<storage, read> a: array<f32>;
|
|
@group(0) @binding(1) var<storage, read> b: array<f32>;
|
|
@group(0) @binding(2) var<storage, read_write> c: array<f32>;
|
|
|
|
struct AddParams {
|
|
size: u32,
|
|
}
|
|
@group(0) @binding(3) var<uniform> params: AddParams;
|
|
|
|
@compute @workgroup_size(256)
|
|
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
|
|
let idx = global_id.x;
|
|
|
|
if (idx >= params.size) {
|
|
return;
|
|
}
|
|
|
|
c[idx] = a[idx] + b[idx];
|
|
}
|
|
";
|
|
|
|
/// Collection of pre-built shaders for inference
|
|
#[wasm_bindgen]
|
|
pub struct WebGPUShaders;
|
|
|
|
#[wasm_bindgen]
|
|
impl WebGPUShaders {
|
|
/// Get GEMM shader source
|
|
#[wasm_bindgen]
|
|
pub fn gemm() -> String {
|
|
GEMM_SHADER.to_string()
|
|
}
|
|
|
|
/// Get softmax shader source
|
|
#[wasm_bindgen]
|
|
pub fn softmax() -> String {
|
|
SOFTMAX_SHADER.to_string()
|
|
}
|
|
|
|
/// Get layer norm shader source
|
|
#[wasm_bindgen]
|
|
pub fn layer_norm() -> String {
|
|
LAYER_NORM_SHADER.to_string()
|
|
}
|
|
|
|
/// Get GELU shader source
|
|
#[wasm_bindgen]
|
|
pub fn gelu() -> String {
|
|
GELU_SHADER.to_string()
|
|
}
|
|
|
|
/// Get attention shader source
|
|
#[wasm_bindgen]
|
|
pub fn attention() -> String {
|
|
ATTENTION_SHADER.to_string()
|
|
}
|
|
|
|
/// Get RoPE shader source
|
|
#[wasm_bindgen]
|
|
pub fn rope() -> String {
|
|
ROPE_SHADER.to_string()
|
|
}
|
|
|
|
/// Get add shader source
|
|
#[wasm_bindgen]
|
|
pub fn add() -> String {
|
|
ADD_SHADER.to_string()
|
|
}
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
|
|
#[test]
|
|
fn test_webgpu_status() {
|
|
assert_eq!(WebGPUStatus::Available, WebGPUStatus::Available);
|
|
assert_ne!(WebGPUStatus::Available, WebGPUStatus::NotSupported);
|
|
}
|
|
|
|
#[test]
|
|
fn test_limits_default() {
|
|
let limits = WebGPULimits::default();
|
|
assert_eq!(limits.max_buffer_size, 256 * 1024 * 1024);
|
|
assert_eq!(limits.max_compute_workgroup_size_x, 256);
|
|
}
|
|
|
|
#[test]
|
|
fn test_shader_sources() {
|
|
assert!(!GEMM_SHADER.is_empty());
|
|
assert!(!SOFTMAX_SHADER.is_empty());
|
|
assert!(!LAYER_NORM_SHADER.is_empty());
|
|
assert!(!GELU_SHADER.is_empty());
|
|
assert!(!ATTENTION_SHADER.is_empty());
|
|
assert!(!ROPE_SHADER.is_empty());
|
|
assert!(!ADD_SHADER.is_empty());
|
|
|
|
// Verify shaders contain expected elements
|
|
assert!(GEMM_SHADER.contains("@compute"));
|
|
assert!(SOFTMAX_SHADER.contains("@workgroup_size"));
|
|
assert!(ATTENTION_SHADER.contains("@binding"));
|
|
}
|
|
|
|
#[test]
|
|
fn test_buffer_usage() {
|
|
assert_eq!(WebGPUBufferUsage::Storage, WebGPUBufferUsage::Storage);
|
|
assert_ne!(WebGPUBufferUsage::Storage, WebGPUBufferUsage::Uniform);
|
|
}
|
|
}
|