1use std::collections::HashMap;
6use std::sync::{Arc, Mutex};
7
8use wgpu::util::DeviceExt;
9
10use crate::error::{ForgeError, Result};
11
12pub mod ops;
13
14pub 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
70pub 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 #[cfg(not(target_arch = "wasm32"))]
87 pub fn new() -> Result<Arc<Self>> {
88 pollster::block_on(Self::new_async())
89 }
90
91 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 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 #[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 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 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 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 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
347mod 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
397pub 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}