use std::collections::HashMap;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, Mutex};
use wgpu::util::DeviceExt;
use crate::error::{ForgeError, Result};
pub mod ops;
pub const OFFSET_ALIGN_BYTES: usize = 256;
const SHADERS_CORE: &[(&str, &str)] = &[
("add", include_str!("../../../shaders/add.wgsl")),
("gelu", include_str!("../../../shaders/gelu.wgsl")),
("matmul", include_str!("../../../shaders/matmul.wgsl")),
("softmax", include_str!("../../../shaders/softmax.wgsl")),
("layernorm", include_str!("../../../shaders/layernorm.wgsl")),
("embedding", include_str!("../../../shaders/embedding.wgsl")),
(
"split_heads",
include_str!("../../../shaders/split_heads.wgsl"),
),
(
"merge_heads",
include_str!("../../../shaders/merge_heads.wgsl"),
),
("kv_append", include_str!("../../../shaders/kv_append.wgsl")),
];
#[cfg(feature = "train")]
const SHADERS_TRAIN: &[(&str, &str)] = &[
("gelu_bwd", include_str!("../../../shaders/gelu_bwd.wgsl")),
(
"softmax_bwd",
include_str!("../../../shaders/softmax_bwd.wgsl"),
),
(
"layernorm_bwd_dx",
include_str!("../../../shaders/layernorm_bwd_dx.wgsl"),
),
(
"layernorm_bwd_dp",
include_str!("../../../shaders/layernorm_bwd_dp.wgsl"),
),
("sum_rows", include_str!("../../../shaders/sum_rows.wgsl")),
(
"scatter_add",
include_str!("../../../shaders/scatter_add.wgsl"),
),
(
"gather_nll",
include_str!("../../../shaders/gather_nll.wgsl"),
),
("ce_bwd", include_str!("../../../shaders/ce_bwd.wgsl")),
("dropout", include_str!("../../../shaders/dropout.wgsl")),
(
"unsplit_heads",
include_str!("../../../shaders/unsplit_heads.wgsl"),
),
(
"unmerge_heads",
include_str!("../../../shaders/unmerge_heads.wgsl"),
),
("sumsq", include_str!("../../../shaders/sumsq.wgsl")),
("scale", include_str!("../../../shaders/scale.wgsl")),
("adamw", include_str!("../../../shaders/adamw.wgsl")),
];
#[cfg(not(feature = "train"))]
const SHADERS_TRAIN: &[(&str, &str)] = &[];
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct Stats {
pub dispatches: usize,
pub submits: usize,
pub buffers_created: usize,
pub bytes_allocated: usize,
}
impl Stats {
pub fn since(self, earlier: Stats) -> Stats {
Stats {
dispatches: self.dispatches - earlier.dispatches,
submits: self.submits - earlier.submits,
buffers_created: self.buffers_created - earlier.buffers_created,
bytes_allocated: self.bytes_allocated - earlier.bytes_allocated,
}
}
}
#[derive(Default)]
struct Counters {
dispatches: AtomicUsize,
submits: AtomicUsize,
buffers_created: AtomicUsize,
bytes_allocated: AtomicUsize,
}
#[derive(Default)]
struct Pool {
free: HashMap<u64, Vec<wgpu::Buffer>>,
bytes: usize,
}
const MAX_POOL_BYTES: usize = 512 * 1024 * 1024;
pub struct PooledBuffer {
buf: Option<wgpu::Buffer>,
ctx: Arc<WgpuContext>,
size: u64,
recycle: bool,
}
impl std::ops::Deref for PooledBuffer {
type Target = wgpu::Buffer;
fn deref(&self) -> &wgpu::Buffer {
self.buf.as_ref().expect("PooledBuffer used after drop")
}
}
impl std::fmt::Debug for PooledBuffer {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "PooledBuffer({} B)", self.size)
}
}
impl Drop for PooledBuffer {
fn drop(&mut self) {
let Some(buf) = self.buf.take() else { return };
if !self.recycle {
return;
}
let mut pool = self.ctx.pool.lock().unwrap();
if pool.bytes + self.size as usize > MAX_POOL_BYTES {
return; }
pool.bytes += self.size as usize;
pool.free.entry(self.size).or_default().push(buf);
}
}
pub struct DispatchScope {
ctx: Arc<WgpuContext>,
}
impl Drop for DispatchScope {
fn drop(&mut self) {
let outermost = {
let mut s = self.ctx.scope.lock().unwrap();
s.depth -= 1;
s.depth == 0
};
if outermost {
self.ctx.flush();
}
}
}
#[derive(Default)]
struct ScopeState {
encoder: Option<wgpu::CommandEncoder>,
pass: Option<wgpu::ComputePass<'static>>,
depth: usize,
}
type Kernel = (Arc<wgpu::ComputePipeline>, Arc<wgpu::BindGroupLayout>);
pub struct WgpuContext {
pub device: wgpu::Device,
pub queue: wgpu::Queue,
pub adapter_info: wgpu::AdapterInfo,
pipelines: Mutex<HashMap<&'static str, Kernel>>,
counters: Counters,
pool: Mutex<Pool>,
scope: Mutex<ScopeState>,
}
impl std::fmt::Debug for WgpuContext {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "WgpuContext({})", self.adapter_info.name)
}
}
impl WgpuContext {
#[cfg(not(target_arch = "wasm32"))]
pub fn new() -> Result<Arc<Self>> {
pollster::block_on(Self::new_async())
}
pub async fn new_async() -> Result<Arc<Self>> {
let instance = wgpu::Instance::new(&wgpu::InstanceDescriptor::default());
let adapter = instance
.request_adapter(&wgpu::RequestAdapterOptions {
power_preference: wgpu::PowerPreference::HighPerformance,
..Default::default()
})
.await
.map_err(|e| ForgeError::Wgpu(format!("no adapter: {e}")))?;
let adapter_info = adapter.get_info();
let (device, queue) = adapter
.request_device(&wgpu::DeviceDescriptor {
label: Some("forge"),
required_limits: adapter.limits(),
..Default::default()
})
.await
.map_err(|e| ForgeError::Wgpu(format!("request_device: {e}")))?;
Ok(Arc::new(WgpuContext {
device,
queue,
adapter_info,
pipelines: Mutex::new(HashMap::new()),
counters: Counters::default(),
pool: Mutex::new(Pool::default()),
scope: Mutex::new(ScopeState::default()),
}))
}
pub fn scope(self: &Arc<Self>) -> DispatchScope {
let mut s = self.scope.lock().unwrap();
s.depth += 1;
DispatchScope { ctx: self.clone() }
}
pub fn flush(&self) {
let encoder = {
let mut s = self.scope.lock().unwrap();
s.pass = None; s.encoder.take()
};
if let Some(encoder) = encoder {
self.submit(encoder);
}
}
fn submit_with_copies(&self, copies: &[(&wgpu::Buffer, u64, &wgpu::Buffer, u64)]) {
let mut encoder = {
let mut s = self.scope.lock().unwrap();
s.pass = None;
s.encoder.take()
}
.unwrap_or_else(|| self.device.create_command_encoder(&Default::default()));
for (src, src_off, dst, size) in copies {
encoder.copy_buffer_to_buffer(src, *src_off, dst, 0, *size);
}
self.submit(encoder);
}
pub fn create_pooled(self: &Arc<Self>, size_bytes: usize) -> PooledBuffer {
let size = size_bytes.max(4).next_multiple_of(4) as u64;
let recycled = {
let mut pool = self.pool.lock().unwrap();
let hit = pool.free.get_mut(&size).and_then(Vec::pop);
if hit.is_some() {
pool.bytes -= size as usize;
}
hit
};
let buf = recycled.unwrap_or_else(|| self.create_storage(size as usize));
PooledBuffer {
buf: Some(buf),
ctx: self.clone(),
size,
recycle: true,
}
}
pub fn stats(&self) -> Stats {
Stats {
dispatches: self.counters.dispatches.load(Ordering::Relaxed),
submits: self.counters.submits.load(Ordering::Relaxed),
buffers_created: self.counters.buffers_created.load(Ordering::Relaxed),
bytes_allocated: self.counters.bytes_allocated.load(Ordering::Relaxed),
}
}
fn submit(&self, encoder: wgpu::CommandEncoder) {
self.counters.submits.fetch_add(1, Ordering::Relaxed);
self.queue.submit([encoder.finish()]);
}
fn pipeline(&self, name: &'static str) -> Kernel {
let mut cache = self.pipelines.lock().unwrap();
cache
.entry(name)
.or_insert_with(|| {
let src = SHADERS_CORE
.iter()
.chain(SHADERS_TRAIN)
.find(|(n, _)| *n == name)
.unwrap_or_else(|| panic!("unknown shader {name}"))
.1;
let module = self
.device
.create_shader_module(wgpu::ShaderModuleDescriptor {
label: Some(name),
source: wgpu::ShaderSource::Wgsl(src.into()),
});
let pipeline =
self.device
.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
label: Some(name),
layout: None,
module: &module,
entry_point: Some("main"),
compilation_options: Default::default(),
cache: None,
});
let layout = Arc::new(pipeline.get_bind_group_layout(0));
(Arc::new(pipeline), layout)
})
.clone()
}
pub fn create_storage(&self, size_bytes: usize) -> wgpu::Buffer {
let size = size_bytes.max(4) as u64;
self.counters
.buffers_created
.fetch_add(1, Ordering::Relaxed);
self.counters
.bytes_allocated
.fetch_add(size as usize, Ordering::Relaxed);
self.device.create_buffer(&wgpu::BufferDescriptor {
label: None,
size,
usage: wgpu::BufferUsages::STORAGE
| wgpu::BufferUsages::COPY_DST
| wgpu::BufferUsages::COPY_SRC,
mapped_at_creation: false,
})
}
pub fn create_zeroed(self: &Arc<Self>, size_bytes: usize) -> PooledBuffer {
PooledBuffer {
buf: Some(self.create_storage(size_bytes.max(4))),
ctx: self.clone(),
size: size_bytes as u64,
recycle: false,
}
}
pub fn upload(self: &Arc<Self>, bytes: &[u8]) -> PooledBuffer {
let buf = self.create_storage(bytes.len().max(4));
self.queue.write_buffer(&buf, 0, bytes);
PooledBuffer {
buf: Some(buf),
ctx: self.clone(),
size: bytes.len() as u64,
recycle: false,
}
}
fn stage_copy(
&self,
buf: &wgpu::Buffer,
offset_bytes: usize,
size_bytes: usize,
) -> wgpu::Buffer {
let staging = self.device.create_buffer(&wgpu::BufferDescriptor {
label: None,
size: size_bytes as u64,
usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
});
self.submit_with_copies(&[(buf, offset_bytes as u64, &staging, size_bytes as u64)]);
staging
}
#[cfg(not(target_arch = "wasm32"))]
pub fn readback(
&self,
buf: &wgpu::Buffer,
offset_bytes: usize,
size_bytes: usize,
) -> Result<Vec<u8>> {
let staging = self.stage_copy(buf, offset_bytes, size_bytes);
let slice = staging.slice(..);
let (tx, rx) = std::sync::mpsc::channel();
slice.map_async(wgpu::MapMode::Read, move |r| {
let _ = tx.send(r);
});
self.device
.poll(wgpu::PollType::Wait)
.map_err(|e| ForgeError::Wgpu(format!("poll: {e:?}")))?;
rx.recv()
.map_err(|_| ForgeError::Wgpu("map_async callback dropped".into()))?
.map_err(|e| ForgeError::Wgpu(format!("map_async: {e:?}")))?;
let out = slice.get_mapped_range().to_vec();
staging.unmap();
Ok(out)
}
pub async fn readback_async(
&self,
buf: &wgpu::Buffer,
offset_bytes: usize,
size_bytes: usize,
) -> Result<Vec<u8>> {
let staging = self.stage_copy(buf, offset_bytes, size_bytes);
let slice = staging.slice(..);
let (tx, rx) = oneshot::channel();
slice.map_async(wgpu::MapMode::Read, move |r| tx.send(r));
#[cfg(not(target_arch = "wasm32"))]
self.device
.poll(wgpu::PollType::Wait)
.map_err(|e| ForgeError::Wgpu(format!("poll: {e:?}")))?;
#[cfg(target_arch = "wasm32")]
let _ = self.device.poll(wgpu::PollType::Poll);
rx.await
.map_err(|e| ForgeError::Wgpu(format!("map_async: {e:?}")))?;
let out = slice.get_mapped_range().to_vec();
staging.unmap();
Ok(out)
}
pub async fn readback_many_async(
&self,
regions: &[(&wgpu::Buffer, usize, usize)],
) -> Result<Vec<Vec<u8>>> {
if regions.is_empty() {
return Ok(Vec::new());
}
let staging: Vec<wgpu::Buffer> = regions
.iter()
.map(|(_, _, size_bytes)| {
self.device.create_buffer(&wgpu::BufferDescriptor {
label: None,
size: *size_bytes as u64,
usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
})
})
.collect();
let copies: Vec<_> = regions
.iter()
.zip(&staging)
.map(|((buf, off, size), s)| (*buf, *off as u64, s, *size as u64))
.collect();
self.submit_with_copies(&copies);
let waits: Vec<_> = staging
.iter()
.map(|s| {
let (tx, rx) = oneshot::channel();
s.slice(..)
.map_async(wgpu::MapMode::Read, move |r| tx.send(r));
rx
})
.collect();
#[cfg(not(target_arch = "wasm32"))]
self.device
.poll(wgpu::PollType::Wait)
.map_err(|e| ForgeError::Wgpu(format!("poll: {e:?}")))?;
#[cfg(target_arch = "wasm32")]
let _ = self.device.poll(wgpu::PollType::Poll);
let mut out = Vec::with_capacity(regions.len());
for (rx, s) in waits.into_iter().zip(&staging) {
rx.await
.map_err(|e| ForgeError::Wgpu(format!("map_async: {e:?}")))?;
out.push(s.slice(..).get_mapped_range().to_vec());
s.unmap();
}
Ok(out)
}
pub fn dispatch(
&self,
name: &'static str,
params: &[u32],
buffers: &[(&wgpu::Buffer, usize, usize)],
workgroups: (u32, u32, u32),
) {
self.counters.dispatches.fetch_add(1, Ordering::Relaxed);
let (pipeline, layout) = self.pipeline(name);
let params_buf = self
.device
.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some(name),
contents: bytemuck::cast_slice(params),
usage: wgpu::BufferUsages::UNIFORM,
});
let mut entries = vec![wgpu::BindGroupEntry {
binding: 0,
resource: params_buf.as_entire_binding(),
}];
for (i, (buf, off, size)) in buffers.iter().enumerate() {
debug_assert!(off % OFFSET_ALIGN_BYTES == 0, "storage offset misaligned");
entries.push(wgpu::BindGroupEntry {
binding: (i + 1) as u32,
resource: wgpu::BindingResource::Buffer(wgpu::BufferBinding {
buffer: buf,
offset: *off as u64,
size: Some(std::num::NonZeroU64::new((*size).max(4) as u64).unwrap()),
}),
});
}
let bind = self.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some(name),
layout: &layout,
entries: &entries,
});
let mut scope = self.scope.lock().unwrap();
if scope.depth > 0 {
let ScopeState { encoder, pass, .. } = &mut *scope;
let encoder = encoder
.get_or_insert_with(|| self.device.create_command_encoder(&Default::default()));
let pass = match pass {
Some(p) => p,
None => pass.insert(
encoder
.begin_compute_pass(&Default::default())
.forget_lifetime(),
),
};
pass.set_pipeline(&pipeline);
pass.set_bind_group(0, &bind, &[]);
pass.dispatch_workgroups(workgroups.0, workgroups.1, workgroups.2);
return;
}
drop(scope);
let mut encoder = self.device.create_command_encoder(&Default::default());
{
let mut pass = encoder.begin_compute_pass(&Default::default());
pass.set_pipeline(&pipeline);
pass.set_bind_group(0, &bind, &[]);
pass.dispatch_workgroups(workgroups.0, workgroups.1, workgroups.2);
}
self.submit(encoder);
}
}
mod oneshot {
use std::future::Future;
use std::pin::Pin;
use std::sync::{Arc, Mutex};
use std::task::{Context, Poll, Waker};
struct State<T> {
value: Option<T>,
waker: Option<Waker>,
}
pub struct Sender<T>(Arc<Mutex<State<T>>>);
pub struct Receiver<T>(Arc<Mutex<State<T>>>);
pub fn channel<T>() -> (Sender<T>, Receiver<T>) {
let shared = Arc::new(Mutex::new(State {
value: None,
waker: None,
}));
(Sender(shared.clone()), Receiver(shared))
}
impl<T> Sender<T> {
pub fn send(self, value: T) {
let mut s = self.0.lock().unwrap();
s.value = Some(value);
if let Some(w) = s.waker.take() {
w.wake();
}
}
}
impl<T> Future for Receiver<T> {
type Output = T;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<T> {
let mut s = self.0.lock().unwrap();
match s.value.take() {
Some(v) => Poll::Ready(v),
None => {
s.waker = Some(cx.waker().clone());
Poll::Pending
}
}
}
}
}
pub fn linear_grid(numel: usize) -> (u32, u32, u32) {
let groups = numel.div_ceil(256).max(1) as u32;
if groups <= 65535 {
(groups, 1, 1)
} else {
let y = groups.div_ceil(65535);
(65535, y, 1)
}
}