Skip to main content

forge/backend/wgpu/
mod.rs

1//! WebGPU backend: device context, buffer management, pipeline cache, and
2//! kernel dispatch. Production backend of Forge; the CPU backend is the
3//! numerical reference.
4
5use std::collections::HashMap;
6use std::sync::{Arc, Mutex};
7
8use wgpu::util::DeviceExt;
9
10use crate::error::{ForgeError, Result};
11
12pub mod ops;
13
14/// Storage-buffer offsets must respect this alignment when creating views.
15pub const OFFSET_ALIGN_BYTES: usize = 256;
16
17const SHADERS: &[(&str, &str)] = &[
18    ("add", include_str!("../../../shaders/add.wgsl")),
19    ("gelu", include_str!("../../../shaders/gelu.wgsl")),
20    ("matmul", include_str!("../../../shaders/matmul.wgsl")),
21    ("softmax", include_str!("../../../shaders/softmax.wgsl")),
22    ("layernorm", include_str!("../../../shaders/layernorm.wgsl")),
23    ("embedding", include_str!("../../../shaders/embedding.wgsl")),
24    (
25        "split_heads",
26        include_str!("../../../shaders/split_heads.wgsl"),
27    ),
28    (
29        "merge_heads",
30        include_str!("../../../shaders/merge_heads.wgsl"),
31    ),
32    ("kv_append", include_str!("../../../shaders/kv_append.wgsl")),
33    ("gelu_bwd", include_str!("../../../shaders/gelu_bwd.wgsl")),
34    (
35        "softmax_bwd",
36        include_str!("../../../shaders/softmax_bwd.wgsl"),
37    ),
38    (
39        "layernorm_bwd_dx",
40        include_str!("../../../shaders/layernorm_bwd_dx.wgsl"),
41    ),
42    (
43        "layernorm_bwd_dp",
44        include_str!("../../../shaders/layernorm_bwd_dp.wgsl"),
45    ),
46    ("sum_rows", include_str!("../../../shaders/sum_rows.wgsl")),
47    (
48        "scatter_add",
49        include_str!("../../../shaders/scatter_add.wgsl"),
50    ),
51    (
52        "gather_nll",
53        include_str!("../../../shaders/gather_nll.wgsl"),
54    ),
55    ("ce_bwd", include_str!("../../../shaders/ce_bwd.wgsl")),
56    ("dropout", include_str!("../../../shaders/dropout.wgsl")),
57    (
58        "unsplit_heads",
59        include_str!("../../../shaders/unsplit_heads.wgsl"),
60    ),
61    (
62        "unmerge_heads",
63        include_str!("../../../shaders/unmerge_heads.wgsl"),
64    ),
65    ("sumsq", include_str!("../../../shaders/sumsq.wgsl")),
66    ("scale", include_str!("../../../shaders/scale.wgsl")),
67    ("adamw", include_str!("../../../shaders/adamw.wgsl")),
68];
69
70/// Owns the wgpu device/queue and a cache of compiled compute pipelines.
71pub struct WgpuContext {
72    pub device: wgpu::Device,
73    pub queue: wgpu::Queue,
74    pub adapter_info: wgpu::AdapterInfo,
75    pipelines: Mutex<HashMap<&'static str, Arc<wgpu::ComputePipeline>>>,
76}
77
78impl std::fmt::Debug for WgpuContext {
79    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
80        write!(f, "WgpuContext({})", self.adapter_info.name)
81    }
82}
83
84impl WgpuContext {
85    /// Sync device creation — native only. On wasm use [`Self::new_async`].
86    #[cfg(not(target_arch = "wasm32"))]
87    pub fn new() -> Result<Arc<Self>> {
88        pollster::block_on(Self::new_async())
89    }
90
91    /// Async device creation (works on native and wasm32; roadmap v4,
92    /// pitfall 14: the async form is primary, the sync API is the facade).
93    pub async fn new_async() -> Result<Arc<Self>> {
94        let instance = wgpu::Instance::new(&wgpu::InstanceDescriptor::default());
95        let adapter = instance
96            .request_adapter(&wgpu::RequestAdapterOptions {
97                power_preference: wgpu::PowerPreference::HighPerformance,
98                ..Default::default()
99            })
100            .await
101            .map_err(|e| ForgeError::Wgpu(format!("no adapter: {e}")))?;
102        let adapter_info = adapter.get_info();
103        // GPT-2's token embedding (~147 MiB) exceeds the 128 MiB default
104        // max_storage_buffer_binding_size, so request the adapter's limits.
105        let (device, queue) = adapter
106            .request_device(&wgpu::DeviceDescriptor {
107                label: Some("forge"),
108                required_limits: adapter.limits(),
109                ..Default::default()
110            })
111            .await
112            .map_err(|e| ForgeError::Wgpu(format!("request_device: {e}")))?;
113        Ok(Arc::new(WgpuContext {
114            device,
115            queue,
116            adapter_info,
117            pipelines: Mutex::new(HashMap::new()),
118        }))
119    }
120
121    fn pipeline(&self, name: &'static str) -> Arc<wgpu::ComputePipeline> {
122        let mut cache = self.pipelines.lock().unwrap();
123        cache
124            .entry(name)
125            .or_insert_with(|| {
126                let src = SHADERS
127                    .iter()
128                    .find(|(n, _)| *n == name)
129                    .unwrap_or_else(|| panic!("unknown shader {name}"))
130                    .1;
131                let module = self
132                    .device
133                    .create_shader_module(wgpu::ShaderModuleDescriptor {
134                        label: Some(name),
135                        source: wgpu::ShaderSource::Wgsl(src.into()),
136                    });
137                Arc::new(
138                    self.device
139                        .create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
140                            label: Some(name),
141                            layout: None,
142                            module: &module,
143                            entry_point: Some("main"),
144                            compilation_options: Default::default(),
145                            cache: None,
146                        }),
147                )
148            })
149            .clone()
150    }
151
152    pub fn create_storage(&self, size_bytes: usize) -> wgpu::Buffer {
153        self.device.create_buffer(&wgpu::BufferDescriptor {
154            label: None,
155            size: size_bytes.max(4) as u64,
156            usage: wgpu::BufferUsages::STORAGE
157                | wgpu::BufferUsages::COPY_DST
158                | wgpu::BufferUsages::COPY_SRC,
159            mapped_at_creation: false,
160        })
161    }
162
163    pub fn upload(&self, bytes: &[u8]) -> wgpu::Buffer {
164        let buf = self.create_storage(bytes.len());
165        self.queue.write_buffer(&buf, 0, bytes);
166        buf
167    }
168
169    fn stage_copy(
170        &self,
171        buf: &wgpu::Buffer,
172        offset_bytes: usize,
173        size_bytes: usize,
174    ) -> wgpu::Buffer {
175        let staging = self.device.create_buffer(&wgpu::BufferDescriptor {
176            label: None,
177            size: size_bytes as u64,
178            usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
179            mapped_at_creation: false,
180        });
181        let mut encoder = self.device.create_command_encoder(&Default::default());
182        encoder.copy_buffer_to_buffer(buf, offset_bytes as u64, &staging, 0, size_bytes as u64);
183        self.queue.submit([encoder.finish()]);
184        staging
185    }
186
187    /// Read `size_bytes` starting at `offset_bytes` back to the host.
188    /// Sync facade — native only (`device.poll(Wait)` cannot exist on wasm).
189    #[cfg(not(target_arch = "wasm32"))]
190    pub fn readback(
191        &self,
192        buf: &wgpu::Buffer,
193        offset_bytes: usize,
194        size_bytes: usize,
195    ) -> Result<Vec<u8>> {
196        let staging = self.stage_copy(buf, offset_bytes, size_bytes);
197        let slice = staging.slice(..);
198        let (tx, rx) = std::sync::mpsc::channel();
199        slice.map_async(wgpu::MapMode::Read, move |r| {
200            let _ = tx.send(r);
201        });
202        self.device
203            .poll(wgpu::PollType::Wait)
204            .map_err(|e| ForgeError::Wgpu(format!("poll: {e:?}")))?;
205        rx.recv()
206            .map_err(|_| ForgeError::Wgpu("map_async callback dropped".into()))?
207            .map_err(|e| ForgeError::Wgpu(format!("map_async: {e:?}")))?;
208        let out = slice.get_mapped_range().to_vec();
209        staging.unmap();
210        Ok(out)
211    }
212
213    /// Async readback — the primary form; on wasm the browser event loop
214    /// drives the mapping.
215    pub async fn readback_async(
216        &self,
217        buf: &wgpu::Buffer,
218        offset_bytes: usize,
219        size_bytes: usize,
220    ) -> Result<Vec<u8>> {
221        let staging = self.stage_copy(buf, offset_bytes, size_bytes);
222        let slice = staging.slice(..);
223        let (tx, rx) = oneshot::channel();
224        slice.map_async(wgpu::MapMode::Read, move |r| tx.send(r));
225        #[cfg(not(target_arch = "wasm32"))]
226        self.device
227            .poll(wgpu::PollType::Wait)
228            .map_err(|e| ForgeError::Wgpu(format!("poll: {e:?}")))?;
229        #[cfg(target_arch = "wasm32")]
230        let _ = self.device.poll(wgpu::PollType::Poll);
231        rx.await
232            .map_err(|e| ForgeError::Wgpu(format!("map_async: {e:?}")))?;
233        let out = slice.get_mapped_range().to_vec();
234        staging.unmap();
235        Ok(out)
236    }
237
238    /// Read several regions back in one submit and one fence wait.
239    ///
240    /// [`WgpuContext::readback_async`] costs a submit and a wait *each*, which
241    /// dominates when a single logical step wants several small tensors: the
242    /// attention probe reads `n_layer + 1` per generated token, and one at a
243    /// time that cost more than the decode itself. Staged into one encoder
244    /// they cost one round trip regardless of how many there are.
245    ///
246    /// Regions are returned in the order given.
247    pub async fn readback_many_async(
248        &self,
249        regions: &[(&wgpu::Buffer, usize, usize)],
250    ) -> Result<Vec<Vec<u8>>> {
251        if regions.is_empty() {
252            return Ok(Vec::new());
253        }
254        let mut encoder = self.device.create_command_encoder(&Default::default());
255        let staging: Vec<wgpu::Buffer> = regions
256            .iter()
257            .map(|(buf, offset_bytes, size_bytes)| {
258                let s = self.device.create_buffer(&wgpu::BufferDescriptor {
259                    label: None,
260                    size: *size_bytes as u64,
261                    usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
262                    mapped_at_creation: false,
263                });
264                encoder.copy_buffer_to_buffer(buf, *offset_bytes as u64, &s, 0, *size_bytes as u64);
265                s
266            })
267            .collect();
268        self.queue.submit([encoder.finish()]);
269
270        // Every map request is issued before anything is awaited, so one poll
271        // services all of them.
272        let waits: Vec<_> = staging
273            .iter()
274            .map(|s| {
275                let (tx, rx) = oneshot::channel();
276                s.slice(..)
277                    .map_async(wgpu::MapMode::Read, move |r| tx.send(r));
278                rx
279            })
280            .collect();
281        #[cfg(not(target_arch = "wasm32"))]
282        self.device
283            .poll(wgpu::PollType::Wait)
284            .map_err(|e| ForgeError::Wgpu(format!("poll: {e:?}")))?;
285        #[cfg(target_arch = "wasm32")]
286        let _ = self.device.poll(wgpu::PollType::Poll);
287
288        let mut out = Vec::with_capacity(regions.len());
289        for (rx, s) in waits.into_iter().zip(&staging) {
290            rx.await
291                .map_err(|e| ForgeError::Wgpu(format!("map_async: {e:?}")))?;
292            out.push(s.slice(..).get_mapped_range().to_vec());
293            s.unmap();
294        }
295        Ok(out)
296    }
297
298    /// Dispatch `name` with binding 0 = `params` (uniform, raw words) and
299    /// bindings 1.. = `buffers` (storage). Each buffer entry is
300    /// (buffer, offset_bytes, size_bytes).
301    pub fn dispatch(
302        &self,
303        name: &'static str,
304        params: &[u32],
305        buffers: &[(&wgpu::Buffer, usize, usize)],
306        workgroups: (u32, u32, u32),
307    ) {
308        let pipeline = self.pipeline(name);
309        let params_buf = self
310            .device
311            .create_buffer_init(&wgpu::util::BufferInitDescriptor {
312                label: Some(name),
313                contents: bytemuck::cast_slice(params),
314                usage: wgpu::BufferUsages::UNIFORM,
315            });
316        let mut entries = vec![wgpu::BindGroupEntry {
317            binding: 0,
318            resource: params_buf.as_entire_binding(),
319        }];
320        for (i, (buf, off, size)) in buffers.iter().enumerate() {
321            debug_assert!(off % OFFSET_ALIGN_BYTES == 0, "storage offset misaligned");
322            entries.push(wgpu::BindGroupEntry {
323                binding: (i + 1) as u32,
324                resource: wgpu::BindingResource::Buffer(wgpu::BufferBinding {
325                    buffer: buf,
326                    offset: *off as u64,
327                    size: Some(std::num::NonZeroU64::new((*size).max(4) as u64).unwrap()),
328                }),
329            });
330        }
331        let bind = self.device.create_bind_group(&wgpu::BindGroupDescriptor {
332            label: Some(name),
333            layout: &pipeline.get_bind_group_layout(0),
334            entries: &entries,
335        });
336        let mut encoder = self.device.create_command_encoder(&Default::default());
337        {
338            let mut pass = encoder.begin_compute_pass(&Default::default());
339            pass.set_pipeline(&pipeline);
340            pass.set_bind_group(0, &bind, &[]);
341            pass.dispatch_workgroups(workgroups.0, workgroups.1, workgroups.2);
342        }
343        self.queue.submit([encoder.finish()]);
344    }
345}
346
347/// Minimal single-value channel whose receiver is a `Future` — lets
348/// `map_async` results be awaited without extra dependencies (wasm has no
349/// blocking receive).
350mod oneshot {
351    use std::future::Future;
352    use std::pin::Pin;
353    use std::sync::{Arc, Mutex};
354    use std::task::{Context, Poll, Waker};
355
356    struct State<T> {
357        value: Option<T>,
358        waker: Option<Waker>,
359    }
360
361    pub struct Sender<T>(Arc<Mutex<State<T>>>);
362    pub struct Receiver<T>(Arc<Mutex<State<T>>>);
363
364    pub fn channel<T>() -> (Sender<T>, Receiver<T>) {
365        let shared = Arc::new(Mutex::new(State {
366            value: None,
367            waker: None,
368        }));
369        (Sender(shared.clone()), Receiver(shared))
370    }
371
372    impl<T> Sender<T> {
373        pub fn send(self, value: T) {
374            let mut s = self.0.lock().unwrap();
375            s.value = Some(value);
376            if let Some(w) = s.waker.take() {
377                w.wake();
378            }
379        }
380    }
381
382    impl<T> Future for Receiver<T> {
383        type Output = T;
384        fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<T> {
385            let mut s = self.0.lock().unwrap();
386            match s.value.take() {
387                Some(v) => Poll::Ready(v),
388                None => {
389                    s.waker = Some(cx.waker().clone());
390                    Poll::Pending
391                }
392            }
393        }
394    }
395}
396
397/// Split a linear element count into a (x, y, 1) workgroup grid of 256-thread
398/// groups, respecting the 65535 per-dimension dispatch limit.
399pub fn linear_grid(numel: usize) -> (u32, u32, u32) {
400    let groups = numel.div_ceil(256).max(1) as u32;
401    if groups <= 65535 {
402        (groups, 1, 1)
403    } else {
404        let y = groups.div_ceil(65535);
405        (65535, y, 1)
406    }
407}