1use cubecl::prelude::{CubeCount, CubeDim, CubeElement, CubeType, Sequence, TensorBinding};
11use cubecl_wgpu::WgpuRuntime;
12use std::fmt;
13use std::sync::Arc;
14
15use crate::{
16 AccessError, AllocationDomainId, AllocationId, AllocationKey, BackendAllocation, BackendId,
17 BackendRuntimeCache, DType, DeviceAccessError, DeviceAccessRequest, DeviceId, DeviceKind,
18 Error, GpuBackendKind, HostAccessError, MemoryKind, Placement, PreparedDeviceAccess,
19 ProviderCapabilities, ProviderReadMapping, ProviderWriteMapping, RootBoundSpan,
20 RootResourceExtent, Tensor, TensorBackend, TensorDeviceTransfer, TensorRank, TensorRead,
21 TensorScalar, TypedTensor, TypedTensorView,
22};
23
24const DEFAULT_CUBE_DIM_X: u32 = 256;
25
26mod apple;
27mod error;
28#[cfg(not(target_family = "wasm"))]
29mod event_domain;
30mod exec_session;
31mod gemm;
32#[doc(hidden)]
33pub mod interop;
34mod kernels;
35mod memory;
36mod runtime;
37mod runtime_adapter;
38mod structural;
39
40pub use apple::{AppleContext, AppleTransferStats};
41pub(crate) use error::{unsupported_dtype, unsupported_operation};
42#[doc(hidden)]
43pub use exec_session::{with_webgpu_exec_session, WebGpuExecSession};
44pub use memory::{download_webgpu_tensor, upload_webgpu_tensor};
45pub use runtime::{webgpu_available, WebGpuRuntime, WebGpuRuntimeIdentity};
46pub use runtime_adapter::{
47 webgpu_runtime_engine_id, webgpu_runtime_engine_registration,
48 webgpu_runtime_engine_registration_with_id, webgpu_runtime_hardware_class,
49};
50
51pub(crate) struct WebGpuBuffer {
54 handle: cubecl_runtime::server::Handle,
55 byte_len: usize,
56 device_ordinal: usize,
57 managed: Option<Arc<cubecl_runtime::storage::ManagedResource<cubecl_wgpu::WgpuResource>>>,
58 allocation_domain: AllocationDomainId,
59 allocation_id: AllocationId,
60}
61
62static NEXT_WEBGPU_ALLOCATION_ID: std::sync::atomic::AtomicU64 =
63 std::sync::atomic::AtomicU64::new(1);
64
65impl std::fmt::Debug for WebGpuBuffer {
66 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
67 f.debug_struct("WebGpuBuffer")
68 .field("byte_len", &self.byte_len)
69 .field("device_ordinal", &self.device_ordinal)
70 .field("allocation_domain", &self.allocation_domain)
71 .field("allocation_id", &self.allocation_id)
72 .finish()
73 }
74}
75
76impl WebGpuBuffer {
77 fn new(
78 handle: cubecl_runtime::server::Handle,
79 byte_len: usize,
80 device_ordinal: usize,
81 allocation_domain: AllocationDomainId,
82 ) -> Self {
83 Self {
84 handle,
85 byte_len,
86 device_ordinal,
87 managed: None,
88 allocation_domain,
89 allocation_id: AllocationId::from_backend_id(
90 NEXT_WEBGPU_ALLOCATION_ID.fetch_add(1, std::sync::atomic::Ordering::Relaxed),
91 ),
92 }
93 }
94
95 fn element_len<T: 'static>(&self) -> usize {
96 let element_size = std::mem::size_of::<T>();
97 debug_assert!(element_size != 0 && self.byte_len.is_multiple_of(element_size));
98 self.byte_len / element_size
99 }
100
101 fn new_for_runtime(
102 rt: &WebGpuRuntime,
103 handle: cubecl_runtime::server::Handle,
104 byte_len: usize,
105 op: &'static str,
106 ) -> crate::Result<Self> {
107 let Some(_domain) = rt.allocation_domain() else {
108 return Ok(Self::new(
109 handle,
110 byte_len,
111 rt.device_ordinal(),
112 rt.allocation_domain_id(),
113 ));
114 };
115 let managed = rt
116 .client()
117 .get_resource(handle.clone())
118 .map_err(|error| crate::Error::backend_source(op, error))?;
119 let allocation_id = AllocationId::from_backend_id(managed.resource().allocation_id());
120 Ok(Self {
121 handle,
122 byte_len,
123 device_ordinal: rt.device_ordinal(),
124 managed: Some(Arc::new(managed)),
125 allocation_domain: rt.allocation_domain_id(),
126 allocation_id,
127 })
128 }
129}
130
131#[derive(Debug)]
133pub(crate) struct WebGpuPreparedAccess {
134 handle: cubecl_runtime::server::Handle,
135 byte_len: usize,
136 device_ordinal: usize,
137}
138
139impl PreparedDeviceAccess for WebGpuPreparedAccess {
140 fn as_any(&self) -> &dyn std::any::Any {
141 self
142 }
143
144 fn into_any(self: Box<Self>) -> Box<dyn std::any::Any> {
145 self
146 }
147}
148
149struct WebGpuReadMapping {
150 guard: cubecl_wgpu::WgpuMappedReadGuard,
151 range: std::ops::Range<usize>,
152}
153
154impl std::ops::Deref for WebGpuReadMapping {
155 type Target = [u8];
156
157 fn deref(&self) -> &Self::Target {
158 &self.guard[self.range.clone()]
159 }
160}
161
162impl AsRef<[u8]> for WebGpuReadMapping {
163 fn as_ref(&self) -> &[u8] {
164 self
165 }
166}
167
168struct WebGpuWriteMapping {
169 guard: cubecl_wgpu::WgpuMappedWriteGuard,
170 bytes: Vec<u8>,
171}
172
173impl std::ops::Deref for WebGpuWriteMapping {
174 type Target = [u8];
175
176 fn deref(&self) -> &Self::Target {
177 &self.bytes
178 }
179}
180
181impl std::ops::DerefMut for WebGpuWriteMapping {
182 fn deref_mut(&mut self) -> &mut Self::Target {
183 &mut self.bytes
184 }
185}
186
187impl AsRef<[u8]> for WebGpuWriteMapping {
188 fn as_ref(&self) -> &[u8] {
189 self
190 }
191}
192
193impl AsMut<[u8]> for WebGpuWriteMapping {
194 fn as_mut(&mut self) -> &mut [u8] {
195 self
196 }
197}
198
199impl Drop for WebGpuWriteMapping {
200 fn drop(&mut self) {
201 self.guard.copy_from_slice(&self.bytes);
202 }
203}
204
205fn provider_dtype_size(dtype: DType) -> usize {
206 match dtype {
207 DType::F32 | DType::I32 => core::mem::size_of::<f32>(),
208 DType::F64 | DType::I64 => core::mem::size_of::<f64>(),
209 DType::Bool => core::mem::size_of::<bool>(),
210 DType::C32 => core::mem::size_of::<num_complex::Complex32>(),
211 DType::C64 => core::mem::size_of::<num_complex::Complex64>(),
212 DType::External(_) => 0,
215 }
216}
217
218fn provider_mapping_range(
219 buffer: &WebGpuBuffer,
220 span: RootBoundSpan,
221 dtype: DType,
222) -> Result<std::ops::Range<usize>, AccessError> {
223 let start = span.byte_offset();
224 let end = start
225 .checked_add(span.byte_len())
226 .ok_or_else(|| AccessError::Provider {
227 message: "WebGPU mapping span overflows".to_owned(),
228 })?;
229 if end > buffer.byte_len {
230 return Err(AccessError::Provider {
231 message: "WebGPU mapping span exceeds the allocation".to_owned(),
232 });
233 }
234 let element_size = provider_dtype_size(dtype);
235 if !start.is_multiple_of(element_size) || !span.byte_len().is_multiple_of(element_size) {
236 return Err(AccessError::Provider {
237 message: "WebGPU mapping span is not element-aligned".to_owned(),
238 });
239 }
240 Ok(start..end)
241}
242
243unsafe impl BackendAllocation for WebGpuBuffer {
248 fn root_extent(&self) -> RootResourceExtent {
249 RootResourceExtent::try_new(
250 AllocationKey::new(self.allocation_domain, self.allocation_id),
251 0,
252 self.byte_len,
253 8,
254 )
255 .expect("WebGPU allocation metadata is constructed with a valid extent")
256 }
257
258 fn provider_kind(&self) -> BackendId {
259 BackendId::WebGpu
260 }
261
262 fn capabilities(&self) -> ProviderCapabilities {
263 if self.managed.is_some() {
264 ProviderCapabilities::host()
265 } else {
266 ProviderCapabilities::none()
267 }
268 }
269
270 fn prepare_device_access(
271 &self,
272 request: DeviceAccessRequest<'_>,
273 ) -> Result<Box<dyn PreparedDeviceAccess>, DeviceAccessError> {
274 if request.allocation_domain() != self.allocation_domain
275 || request.allocation_id() != self.allocation_id
276 {
277 return Err(DeviceAccessError::InvalidRequest {
278 message: "prepared request does not match the WebGPU allocation identity"
279 .to_owned(),
280 });
281 }
282 if request.byte_len() > self.byte_len {
283 return Err(DeviceAccessError::InvalidRequest {
284 message: "prepared request exceeds the WebGPU allocation extent".to_owned(),
285 });
286 }
287 Ok(Box::new(WebGpuPreparedAccess {
288 handle: self.handle.clone(),
289 byte_len: self.byte_len,
290 device_ordinal: self.device_ordinal,
291 }))
292 }
293
294 fn map_read(
295 &self,
296 span: RootBoundSpan,
297 dtype: DType,
298 ) -> Result<ProviderReadMapping<'_>, AccessError> {
299 let managed = self.managed.as_ref().ok_or(AccessError::Unsupported {
300 backend: "cubecl-webgpu",
301 })?;
302 let range = provider_mapping_range(self, span, dtype)?;
303 let guard = managed
304 .resource()
305 .map_read()
306 .map_err(|error| AccessError::Provider {
307 message: error.to_string(),
308 })?;
309 if range.end > guard.len() {
310 return Err(AccessError::Provider {
311 message: "WebGPU host mapping is shorter than the checked root extent".to_owned(),
312 });
313 }
314 Ok(ProviderReadMapping::from_guard(WebGpuReadMapping {
315 guard,
316 range,
317 }))
318 }
319
320 fn map_write(
321 &self,
322 span: RootBoundSpan,
323 dtype: DType,
324 ) -> Result<ProviderWriteMapping<'_>, AccessError> {
325 let managed = self.managed.as_ref().ok_or(AccessError::Unsupported {
326 backend: "cubecl-webgpu",
327 })?;
328 let range = provider_mapping_range(self, span, dtype)?;
329 let guard = managed
330 .resource()
331 .map_write()
332 .map_err(|error| AccessError::Provider {
333 message: error.to_string(),
334 })?;
335 if range.end > guard.len() {
336 return Err(AccessError::Provider {
337 message: "WebGPU host mapping is shorter than the checked root extent".to_owned(),
338 });
339 }
340 let bytes = vec![0_u8; range.len()];
341 Ok(ProviderWriteMapping::from_guard(WebGpuWriteMapping {
342 guard,
343 bytes,
344 }))
345 }
346
347 fn as_any(&self) -> &dyn std::any::Any {
348 self
349 }
350
351 fn as_any_mut(&mut self) -> &mut dyn std::any::Any {
352 self
353 }
354}
355
356pub(super) fn prepared_webgpu_tensor<T: TensorScalar + 'static>(
357 tensor: &TypedTensor<T>,
358 op: &'static str,
359) -> crate::Result<WebGpuPreparedAccess> {
360 let prepared = tensor.prepare_device_read(op)?;
361 prepared
362 .into_any()
363 .downcast::<WebGpuPreparedAccess>()
364 .map(|prepared| *prepared)
365 .map_err(|_| crate::Error::runtime_state(op, "expected a WebGPU prepared allocation"))
366}
367
368impl WebGpuPreparedAccess {
369 pub(crate) const fn device_ordinal(&self) -> usize {
370 self.device_ordinal
371 }
372}
373
374pub(super) fn prepared_webgpu_view<T: TensorScalar + 'static, R: TensorRank>(
375 view: &TypedTensorView<'_, T, R>,
376 op: &'static str,
377) -> crate::Result<WebGpuPreparedAccess> {
378 let prepared = view.prepare_device_read(op)?;
379 prepared
380 .into_any()
381 .downcast::<WebGpuPreparedAccess>()
382 .map(|prepared| *prepared)
383 .map_err(|_| crate::Error::runtime_state(op, "expected a WebGPU prepared allocation"))
384}
385
386fn checked_shape_product(op: &'static str, shape: &[usize]) -> crate::Result<usize> {
387 shape
388 .iter()
389 .try_fold(1usize, |acc, &dim| acc.checked_mul(dim))
390 .ok_or_else(|| {
391 Error::invalid_argument(
392 op,
393 "shape",
394 format!("shape product overflow for shape {shape:?}"),
395 )
396 })
397}
398
399fn cube_count_for_len(len: usize) -> crate::Result<CubeCount> {
400 let cubes = len.div_ceil(DEFAULT_CUBE_DIM_X as usize);
401 let cubes = u32::try_from(cubes).map_err(|_| {
402 Error::invalid_argument(
403 "cube_count_for_len",
404 "length",
405 format!(
406 "1D WebGPU launch for {len} elements requires {cubes} cubes, \
407 which exceeds u32::MAX"
408 ),
409 )
410 })?;
411 Ok(CubeCount::Static(cubes.max(1), 1, 1))
412}
413
414fn cube_dim_1d() -> CubeDim {
415 CubeDim::new_1d(DEFAULT_CUBE_DIM_X)
416}
417
418fn comptime_sequence<T: CubeType + Clone>(values: &[T]) -> Sequence<T> {
419 let mut out = Sequence::new();
420 for value in values {
421 out.push(value.clone());
422 }
423 out
424}
425
426fn typed_tensor_binding_with_layout<T: CubeElement + TensorScalar + Clone>(
427 tensor: &TypedTensor<T>,
428 shape: &[usize],
429 strides: &[usize],
430 op: &'static str,
431) -> crate::Result<TensorBinding<WgpuRuntime>> {
432 if shape.len() != strides.len() {
433 return Err(Error::rank_mismatch(op, shape.len(), strides.len()));
434 }
435 let prepared = prepared_webgpu_tensor(tensor, op)?;
436 let layout_len = checked_shape_product(op, shape)?;
437 if layout_len != tensor.n_elements() {
438 return Err(Error::runtime_state(
439 op,
440 format!(
441 "WebGPU tensor binding layout covers {layout_len} elements, tensor has {}",
442 tensor.n_elements()
443 ),
444 ));
445 }
446
447 let (shape, strides) = if shape.is_empty() {
448 (vec![1], vec![1])
449 } else {
450 (shape.to_vec(), strides.to_vec())
451 };
452
453 Ok(unsafe { TensorBinding::from_raw_parts(prepared.handle, strides.into(), shape.into()) })
457}
458
459pub(super) fn ensure_resident_on_runtime<T: TensorScalar + 'static>(
460 rt: &WebGpuRuntime,
461 tensor: &TypedTensor<T>,
462 op: &'static str,
463) -> crate::Result<()> {
464 let view = tensor.as_view();
465 let expected_allocation_domain = rt.allocation_domain_id();
466 let Some(actual_allocation_domain) = tensor.allocation_domain() else {
467 return Err(Error::runtime_state(
468 op,
469 "expected a WebGPU backend tensor, got host storage",
470 ));
471 };
472 if actual_allocation_domain != expected_allocation_domain {
473 return Err(Error::host_access(
474 op,
475 HostAccessError::ForeignDomain {
476 expected: expected_allocation_domain,
477 actual: actual_allocation_domain,
478 },
479 ));
480 }
481 if !matches!(view.backend_family(), Some("webgpu" | "cubecl-webgpu")) {
482 return Err(Error::runtime_state(
483 op,
484 "expected a WebGPU allocation from the selected provider",
485 ));
486 }
487 ensure_placement_resident_on_runtime(rt, tensor.placement(), op)
488}
489
490fn ensure_placement_resident_on_runtime(
491 rt: &WebGpuRuntime,
492 placement: &Placement,
493 op: &'static str,
494) -> crate::Result<()> {
495 let expected_memory = if rt.allocation_domain().is_some() {
496 MemoryKind::Managed
497 } else {
498 MemoryKind::Device
499 };
500 if placement.memory_kind != expected_memory {
501 return Err(Error::runtime_state(
502 op,
503 format!(
504 "expected WebGPU tensor placement, got {:?}",
505 placement.memory_kind
506 ),
507 ));
508 }
509 match &placement.device {
510 Some(device)
511 if device.kind == DeviceKind::Gpu(GpuBackendKind::WebGpu)
512 && device.ordinal == rt.device_ordinal() =>
513 {
514 Ok(())
515 }
516 Some(device) => Err(Error::runtime_state(
517 op,
518 format!(
519 "expected WebGPU tensor resident on webgpu:{}, got {:?}:{}",
520 rt.device_ordinal(),
521 device.kind,
522 device.ordinal
523 ),
524 )),
525 None => Err(Error::runtime_state(
526 op,
527 format!(
528 "expected WebGPU tensor resident on webgpu:{}, got missing device metadata",
529 rt.device_ordinal()
530 ),
531 )),
532 }
533}
534
535pub(super) fn typed_from_webgpu<T: TensorScalar + Send + Sync + 'static>(
536 shape: Vec<usize>,
537 buffer: WebGpuBuffer,
538 rt: &WebGpuRuntime,
539) -> crate::Result<TypedTensor<T>> {
540 let expected_len = checked_shape_product("typed_from_webgpu", &shape)?;
541 if expected_len != buffer.element_len::<T>() {
542 return Err(Error::runtime_state(
543 "typed_from_webgpu",
544 format!(
545 "WebGPU allocation has {} elements, shape requires {expected_len}",
546 buffer.element_len::<T>()
547 ),
548 ));
549 }
550 TypedTensor::from_backend_allocation(shape, Box::new(buffer), webgpu_placement(rt))
551}
552
553fn alloc_output<T: CubeElement + TensorScalar + Clone + Send + Sync + 'static>(
554 rt: &WebGpuRuntime,
555 shape: &[usize],
556 op: &'static str,
557) -> crate::Result<TypedTensor<T>> {
558 let len = checked_shape_product(op, shape)?;
559 let bytes = len.checked_mul(core::mem::size_of::<T>()).ok_or_else(|| {
560 Error::invalid_argument(
561 op,
562 "shape",
563 format!("WebGPU output byte length overflow for shape {shape:?}"),
564 )
565 })?;
566 let handle = rt.client().empty(bytes);
567 let buffer = WebGpuBuffer::new_for_runtime(rt, handle, bytes, op)?;
568 typed_from_webgpu(shape.to_vec(), buffer, rt)
569}
570
571pub(super) fn alloc_tensor_in_runtime(
572 rt: &WebGpuRuntime,
573 dtype: DType,
574 shape: &[usize],
575) -> crate::Result<Tensor> {
576 match dtype {
577 DType::External(_) => Err(Error::unsupported(
580 "apple_alloc",
581 "an externally defined scalar has no WebGPU buffer",
582 )),
583 DType::F32 => alloc_output::<f32>(rt, shape, "apple_alloc").map(Tensor::from_typed::<f32>),
584 DType::F64 => alloc_output::<f64>(rt, shape, "apple_alloc").map(Tensor::from_typed::<f64>),
585 DType::I32 => alloc_output::<i32>(rt, shape, "apple_alloc").map(Tensor::from_typed::<i32>),
586 DType::I64 => alloc_output::<i64>(rt, shape, "apple_alloc").map(Tensor::from_typed::<i64>),
587 DType::C32 => alloc_output::<num_complex::Complex32>(rt, shape, "apple_alloc")
588 .map(Tensor::from_typed::<tenferro_tensor::Complex32>),
589 DType::C64 => alloc_output::<num_complex::Complex64>(rt, shape, "apple_alloc")
590 .map(Tensor::from_typed::<tenferro_tensor::Complex64>),
591 DType::Bool => {
592 let len = checked_shape_product("apple_alloc", shape)?;
593 let handle = rt.client().empty(len);
594 let buffer = WebGpuBuffer::new_for_runtime(rt, handle, len, "apple_alloc")?;
595 Ok(Tensor::from_typed::<bool>(
596 TypedTensor::from_backend_allocation(
597 shape.to_vec(),
598 Box::new(buffer),
599 webgpu_placement(rt),
600 )?,
601 ))
602 }
603 }
604}
605
606fn webgpu_placement(rt: &WebGpuRuntime) -> Placement {
607 Placement {
608 memory_kind: if rt.allocation_domain().is_some() {
609 MemoryKind::Managed
610 } else {
611 MemoryKind::Device
612 },
613 device: Some(DeviceId {
614 kind: DeviceKind::Gpu(GpuBackendKind::WebGpu),
615 ordinal: rt.device_ordinal(),
616 }),
617 cpu_affinity: None,
618 }
619}
620
621#[derive(Clone)]
639pub struct WebGpuBackend {
640 runtime: WebGpuRuntime,
641}
642
643impl fmt::Debug for WebGpuBackend {
644 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
645 f.debug_struct("WebGpuBackend")
646 .field("runtime", &self.runtime)
647 .finish_non_exhaustive()
648 }
649}
650
651impl WebGpuBackend {
652 pub fn new(device_ordinal: usize) -> crate::Result<Self> {
668 WebGpuRuntime::new(device_ordinal).map(Self::from_runtime)
669 }
670
671 pub fn new_default() -> crate::Result<Self> {
687 WebGpuRuntime::new_default().map(Self::from_runtime)
688 }
689
690 pub fn from_runtime(runtime: WebGpuRuntime) -> Self {
700 Self { runtime }
701 }
702
703 pub fn runtime(&self) -> &WebGpuRuntime {
713 &self.runtime
714 }
715
716 pub fn runtime_identity(&self) -> WebGpuRuntimeIdentity {
730 self.runtime.runtime_identity()
731 }
732
733 pub fn synchronize(&self) -> crate::Result<()> {
749 self.runtime.synchronize()
750 }
751}
752
753pub(crate) fn unsupported_op(op: &'static str) -> crate::Error {
754 crate::Error::unsupported(
755 op,
756 "WebGPU backend does not support this operation yet; upload/download explicitly and use a supported backend operation",
757 )
758}
759
760macro_rules! unsupported {
761 ($op:literal) => {
762 Err(unsupported_op($op))
763 };
764}
765pub(crate) use unsupported;
766
767impl TensorDeviceTransfer for WebGpuBackend {
768 fn download_to_host(&mut self, tensor: TensorRead<'_>) -> crate::Result<Tensor> {
769 let tensor = tensor.as_tensor().ok_or_else(|| {
770 crate::Error::unsupported(
771 "WebGpuBackend::download_to_host",
772 "WebGPU transfer currently requires an owned tensor; materialize a view explicitly first",
773 )
774 })?;
775 download_webgpu_tensor(self.runtime(), tensor)
776 }
777
778 fn upload_host_tensor(&mut self, tensor: TensorRead<'_>) -> crate::Result<Tensor> {
779 let tensor = tensor.as_tensor().ok_or_else(|| {
780 crate::Error::unsupported(
781 "WebGpuBackend::upload_host_tensor",
782 "WebGPU transfer currently requires an owned tensor; materialize a view explicitly first",
783 )
784 })?;
785 upload_webgpu_tensor(self.runtime(), tensor)
786 }
787}
788
789impl BackendRuntimeCache for WebGpuBackend {
790 type RuntimeCache = ();
791}
792
793impl TensorBackend for WebGpuBackend {}