use super::{BackendError, Device, DeviceId, DeviceInfo, ErrorStatus, Event, MemoryPool, PoolId};
use crate::{
DType,
backend::{DTypeCapability, DeviceProgramId, PoolBufferId},
kernel::{IdxScope, Kernel, MemScope, Op},
shape::Dim,
slab::Slab,
};
use nanoserde::DeJson;
use pollster::FutureExt;
use std::{sync::Arc, time::Duration};
use wgpu::{
BindGroupLayout, BufferDescriptor, BufferUsages, ComputePipeline, PowerPreference, ShaderModule, SubmissionIndex,
wgt::PollType,
};
#[derive(DeJson, Debug)]
#[nserde(default)]
pub struct WGPUConfig {
enabled: bool,
}
impl Default for WGPUConfig {
fn default() -> Self {
Self { enabled: true }
}
}
#[derive(Debug)]
pub struct WGPUMemoryPool {
free_bytes: Dim,
device: Arc<wgpu::Device>,
queue: Arc<wgpu::Queue>,
buffers: Slab<PoolBufferId, wgpu::Buffer>,
}
#[derive(Debug)]
pub struct WGPUDevice {
dev_info: DeviceInfo,
memory_pool_id: PoolId,
device: Arc<wgpu::Device>,
#[allow(unused)]
adapter: wgpu::Adapter,
programs: Slab<DeviceProgramId, WGPUProgram>,
queue: Arc<wgpu::Queue>,
}
#[derive(Debug, Clone)]
pub struct WGPUEvent {
submission_index: Option<SubmissionIndex>,
}
#[derive(Debug)]
#[allow(dead_code)]
pub(super) struct WGPUProgram {
name: String,
gws: Vec<u64>,
arg_ro_flags: Vec<bool>,
shader: ShaderModule,
pipeline: ComputePipeline,
bind_group_layout: BindGroupLayout,
}
pub(super) fn initialize_device(
config: &WGPUConfig,
memory_pools: &mut Slab<PoolId, MemoryPool>,
devices: &mut Slab<DeviceId, Device>,
debug_dev: bool,
) -> Result<(), BackendError> {
if !config.enabled {
if debug_dev {
println!("[WGPU] configured out");
}
return Ok(());
}
let power_preference = PowerPreference::from_env().unwrap_or(wgpu::PowerPreference::HighPerformance);
let instance = wgpu::Instance::new(wgpu::InstanceDescriptor {
backends: wgpu::Backends::all(),
flags: wgpu::InstanceFlags::empty(),
memory_budget_thresholds: wgpu::MemoryBudgetThresholds { for_resource_creation: None, for_device_loss: None },
backend_options: wgpu::BackendOptions::from_env_or_default(),
display: None,
});
if debug_dev {
println!("[WGPU] requesting device with {power_preference:#?} power preference");
}
let (wgpu_adapter, wgpu_device, wgpu_queue, info) = async {
let adapter = instance
.request_adapter(&wgpu::RequestAdapterOptions { power_preference, ..Default::default() })
.await
.expect("Failed at adapter creation.");
let info = adapter.get_info();
let mut features = wgpu::Features::empty();
if adapter.features().contains(wgpu::Features::SHADER_F64) {
features |= wgpu::Features::SHADER_F64;
}
if adapter.features().contains(wgpu::Features::SHADER_INT64) {
features |= wgpu::Features::SHADER_INT64;
}
if adapter.features().contains(wgpu::Features::SHADER_F16) {
features |= wgpu::Features::SHADER_F16;
}
let (device, queue) = adapter
.request_device(&wgpu::DeviceDescriptor {
label: None,
required_features: features,
required_limits: adapter.limits(),
experimental_features: wgpu::ExperimentalFeatures::disabled(),
memory_hints: wgpu::MemoryHints::default(),
trace: wgpu::Trace::Off,
})
.await
.expect("Failed at device creation");
(adapter, device, queue, info)
}
.block_on();
if debug_dev {
println!("[WGPU] {} ({}) — {:#?}", info.name, info.device, info.backend);
}
let device = Arc::new(wgpu_device);
let queue = Arc::new(wgpu_queue);
let pool = MemoryPool::WGPU(WGPUMemoryPool {
free_bytes: 1_000_000_000,
device: device.clone(),
queue: queue.clone(),
buffers: Slab::new(),
});
if debug_dev {
println!("[WGPU] device total memory: {} MB", 1_000_000_000u64 / (1024 * 1024));
}
memory_pools.push(pool);
let limits = device.limits();
let wgpu_features = wgpu_adapter.features();
let dtype_capability = {
let mut ops = [DTypeCapability::all(); DType::N_DTYPES];
ops[DType::F64 as usize] = DTypeCapability::none();
if !wgpu_features.contains(wgpu::Features::SHADER_INT64) {
ops[DType::I64 as usize] = DTypeCapability::none();
ops[DType::U64 as usize] = DTypeCapability::none();
}
if !wgpu_features.contains(wgpu::Features::SHADER_F16) {
ops[DType::F16 as usize] = DTypeCapability::none();
}
ops[DType::BF16 as usize] = DTypeCapability::none();
ops[DType::U8 as usize] = DTypeCapability::none();
ops[DType::I8 as usize] = DTypeCapability::none();
ops[DType::U16 as usize] = DTypeCapability::none();
ops[DType::I16 as usize] = DTypeCapability::none();
ops
};
devices.push(Device::WGPU(WGPUDevice {
device,
adapter: wgpu_adapter,
dev_info: DeviceInfo {
compute: 1024 * 1024 * 1024 * 1024,
max_global_work_dims: vec![100_000; 3],
max_local_threads: Dim::from(limits.max_compute_invocations_per_workgroup),
max_local_work_dims: vec![
Dim::from(limits.max_compute_workgroup_size_x),
Dim::from(limits.max_compute_workgroup_size_y),
Dim::from(limits.max_compute_workgroup_size_z),
],
preferred_vector_size: 4,
local_mem_size: 64 * 1024,
max_register_bytes: 512,
tensor_cores: false,
warp_size: 32,
has_native_exp2: true,
supported_vec_lens: vec![2, 3, 4],
dtype_capability,
},
memory_pool_id: PoolId::from(usize::from(memory_pools.len()) - 1),
programs: Slab::new(),
queue,
}));
Ok(())
}
impl WGPUMemoryPool {
#[allow(clippy::unused_self)]
pub const fn deinitialize(&mut self) {}
pub const fn free_bytes(&self) -> Dim {
self.free_bytes
}
pub fn allocate(&mut self, bytes: Dim) -> Result<(PoolBufferId, Event), BackendError> {
const ALIGN: Dim = wgpu::COPY_BUFFER_ALIGNMENT;
let bytes = bytes.div_ceil(ALIGN) * ALIGN;
if bytes > self.free_bytes {
return Err(BackendError { status: ErrorStatus::MemoryAllocation, context: "".into() });
}
let buffer = self.device.create_buffer(&BufferDescriptor {
label: None,
size: bytes as u64,
usage: BufferUsages::from_bits_truncate(
BufferUsages::STORAGE.bits() | BufferUsages::COPY_SRC.bits() | BufferUsages::COPY_DST.bits(),
),
mapped_at_creation: false,
});
let id = self.buffers.push(buffer);
let event = Event::WGPU(WGPUEvent { submission_index: None });
Ok((id, event))
}
pub fn deallocate(&mut self, buffer_id: PoolBufferId, event_wait_list: Vec<Event>) {
drop(event_wait_list);
let buffer = unsafe { self.buffers.remove_and_return(buffer_id) };
buffer.destroy();
}
#[allow(clippy::unnecessary_wraps)]
pub fn host_to_pool(
&mut self,
src: &[u8],
dst: PoolBufferId,
event_wait_list: Vec<Event>,
) -> Result<super::Event, BackendError> {
const ALIGN: usize = wgpu::COPY_BUFFER_ALIGNMENT as usize;
drop(event_wait_list);
let dst = &self.buffers[dst];
let aligned_len = src.len().div_ceil(ALIGN);
if aligned_len > src.len() {
let mut padded: [u8; ALIGN] = [0; ALIGN];
let full_chunks = src.len() / ALIGN;
let remaining = src.len() % ALIGN;
if full_chunks > 0 {
self.queue.write_buffer(dst, 0, &src[..full_chunks * ALIGN]);
}
if remaining > 0 {
padded[..remaining].copy_from_slice(&src[full_chunks * ALIGN..]);
self.queue.write_buffer(dst, (full_chunks * ALIGN) as u64, &padded);
}
} else {
self.queue.write_buffer(dst, 0, src);
}
let encoder = self.device.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: Some("GpuBuffer::write") });
self.queue.submit(Some(encoder.finish()));
Ok(Event::WGPU(WGPUEvent { submission_index: None }))
}
#[allow(clippy::unnecessary_box_returns)]
#[allow(clippy::unnecessary_wraps)]
pub fn pool_to_host(&mut self, src: PoolBufferId, dst: &mut [u8], event_wait_list: Vec<Event>) -> Result<(), BackendError> {
drop(event_wait_list);
let src = &self.buffers[src];
let download_buffer = self.device.create_buffer(&wgpu::BufferDescriptor {
label: Some("DownloadBuffer"), size: dst.len() as u64,
usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST, mapped_at_creation: false,
});
let mut encoder =
self.device.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: Some("CopyBufferEncoder") });
encoder.copy_buffer_to_buffer(
src,
0, &download_buffer,
0, dst.len() as u64, );
let command_buffer = encoder.finish();
self.queue.submit(Some(command_buffer));
let (tx, rx) = std::sync::mpsc::channel();
download_buffer.map_async(wgpu::MapMode::Read, 0..download_buffer.size(), move |result| {
tx.send(result).unwrap();
});
self.device.poll(wgpu::PollType::Wait { submission_index: None, timeout: None }).unwrap();
let mapping_result = rx.recv().unwrap();
mapping_result.unwrap();
let mapped_range = download_buffer.get_mapped_range(0..download_buffer.size());
dst.copy_from_slice(&mapped_range);
drop(mapped_range); download_buffer.unmap();
Ok(())
}
#[allow(clippy::unnecessary_box_returns)]
#[allow(clippy::unnecessary_wraps)]
pub fn sync_events(&mut self, events: Vec<Event>) -> Result<(), BackendError> {
for event in events {
if let Event::WGPU(event) = event {
_ = self
.device
.poll(PollType::Wait { submission_index: event.submission_index, timeout: Some(Duration::from_mins(5)) });
}
}
Ok(())
}
pub fn pool_to_pool(
&mut self,
src_pool: &mut MemoryPool,
src: PoolBufferId,
dst: PoolBufferId,
event_wait_list: Vec<Event>,
) -> Result<Event, BackendError> {
match src_pool {
MemoryPool::Host(src_pool) => self.host_to_pool(src_pool.get_buffer(src), dst, event_wait_list),
_ => todo!("pool_to_pool from {:?} to WGPU", std::mem::discriminant(src_pool)),
}
}
pub fn release_events(&mut self, events: Vec<Event>) {
drop(events);
}
}
impl WGPUDevice {
#[allow(clippy::unused_self)]
pub const fn deinitialize(&mut self) {}
pub const fn info(&self) -> &DeviceInfo {
&self.dev_info
}
pub const fn memory_pool_id(&self) -> PoolId {
self.memory_pool_id
}
pub const fn free_compute(&self) -> u128 {
self.dev_info.compute
}
pub fn compile(&mut self, kernel: &Kernel, debug_asm: bool) -> Result<DeviceProgramId, BackendError> {
let mut gws = vec![Dim::from(1u64); 3];
let mut lws = vec![Dim::from(1u64); 3];
let mut op_id = kernel.head;
while !op_id.is_null() {
match kernel.ops[op_id].op {
Op::Index { len, axis, scope } => match scope {
IdxScope::Group => gws[axis as usize] = len,
IdxScope::Local => lws[axis as usize] = len,
IdxScope::Warp => todo!(),
},
_ => {}
}
op_id = kernel.next_op(op_id);
}
let spirv_words = kernel.generate_spirv(debug_asm)?;
let shader_module = self.device.create_shader_module(wgpu::ShaderModuleDescriptor {
label: None,
source: wgpu::ShaderSource::SpirV(std::borrow::Cow::Owned(spirv_words)),
});
if lws.iter().product::<u64>() > self.dev_info.max_local_threads as u64 {
return Err(BackendError { status: ErrorStatus::KernelCompilation, context: "Invalid local work size.".into() });
}
let name = format!(
"k_{}__{}",
gws.iter().map(ToString::to_string).collect::<Vec<_>>().join("_"),
lws.iter().map(ToString::to_string).collect::<Vec<_>>().join("_"),
);
let mut arg_ro_flags = Vec::new();
let mut op_id = kernel.head;
while !op_id.is_null() {
if let &Op::Define { dtype: _, scope, ro, len: _ } = kernel.at(op_id) {
if scope == MemScope::Global {
arg_ro_flags.push(ro);
}
}
op_id = kernel.next_op(op_id);
}
let bg_layout_entries: Vec<wgpu::BindGroupLayoutEntry> = arg_ro_flags
.iter()
.enumerate()
.map(|(bind_id, ro)| wgpu::BindGroupLayoutEntry {
binding: u32::try_from(bind_id).unwrap(),
visibility: wgpu::ShaderStages::COMPUTE,
ty: wgpu::BindingType::Buffer {
has_dynamic_offset: false,
min_binding_size: None,
ty: wgpu::BufferBindingType::Storage { read_only: *ro },
},
count: None,
})
.collect();
let bind_group_layout =
self.device.create_bind_group_layout(&wgpu::BindGroupLayoutDescriptor { label: None, entries: &bg_layout_entries });
let pipeline_layout = self.device.create_pipeline_layout(&wgpu::PipelineLayoutDescriptor {
label: None,
bind_group_layouts: &[Some(&bind_group_layout)],
immediate_size: 0,
});
let pipeline = self.device.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
label: None,
module: &shader_module,
entry_point: Some(&name),
layout: Some(&pipeline_layout),
cache: None,
compilation_options: wgpu::PipelineCompilationOptions::default(),
});
let id = self.programs.push(WGPUProgram { name, gws, arg_ro_flags, shader: shader_module, pipeline, bind_group_layout });
Ok(id)
}
pub fn release(&mut self, program_id: DeviceProgramId) {
self.programs.remove(program_id);
}
#[allow(clippy::unnecessary_wraps)]
pub fn launch(
&mut self,
program_id: DeviceProgramId,
memory_pool: &mut WGPUMemoryPool,
args: &[PoolBufferId],
event_wait_list: Vec<Event>,
) -> Result<Event, BackendError> {
drop(event_wait_list);
let program = &self.programs[program_id];
let binds: Vec<wgpu::BindGroupEntry> = args
.iter()
.enumerate()
.map(|(bind_id, &arg)| {
let buffer = &memory_pool.buffers[arg];
wgpu::BindGroupEntry { binding: u32::try_from(bind_id).unwrap(), resource: buffer.as_entire_binding() }
})
.collect();
let set = self.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &program.bind_group_layout,
entries: &binds,
});
let mut encoder = self.device.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: Some("Kernel::enqueue") });
{
let mut cpass = encoder
.begin_compute_pass(&wgpu::ComputePassDescriptor { label: Some("Kernel::enqueue"), timestamp_writes: None });
cpass.set_pipeline(&program.pipeline);
cpass.set_bind_group(0, &set, &[]);
cpass.insert_debug_marker(&program.name);
cpass.dispatch_workgroups(
u32::try_from(program.gws.first().copied().unwrap_or(1)).unwrap(),
u32::try_from(program.gws.get(1).copied().unwrap_or(1)).unwrap(),
u32::try_from(program.gws.get(2).copied().unwrap_or(1)).unwrap(),
);
}
let submission_index = Some(self.queue.submit(Some(encoder.finish())));
Ok(Event::WGPU(WGPUEvent { submission_index }))
}
}