Skip to main content

tenferro_gpu/webgpu/
mod.rs

1//! CubeCL WebGPU provider runtime and backend skeleton.
2//!
3//! The WebGPU/Metal backend is experimental and implements a narrow operation
4//! subset: explicit transfers, `F32`/`C32` `dot_general` (and the einsum paths
5//! that lower to it), `F32`/`I32` transpose and compaction, and the Apple Metal
6//! FFT. Elementwise math, reductions, reshape/broadcast, indexing and linear
7//! algebra return typed unsupported errors instead of falling back to CPU. The
8//! devices-and-gpu guide compares this subset with CUDA.
9
10use 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
51/// Scalar-independent WebGPU allocation stored behind tensor backend-buffer
52/// trait objects; dtype is carried by the borrowed tensor descriptor.
53pub(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/// Opaque provider state produced by the shared storage root.
132#[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        // INVARIANT: WebGPU provider buffers are sized for the preset scalars the
213        // provider supports. An externally defined scalar has no fixed width.
214        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
243// SAFETY: WebGpuBuffer owns exactly one CubeCL allocation handle. Its provider
244// guards retain the underlying managed resource for every borrowed mapping;
245// the root importer consumes the buffer exactly once and the provider handle
246// is never used as a public ownership authority.
247unsafe 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    // SAFETY: The tensor root prepared the provider allocation for this exact
454    // descriptor before the binding is constructed. The caller-provided
455    // shape/stride metadata covers the validated logical tensor extent.
456    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        // An externally defined scalar has no WebGPU buffer mapping, so the
578        // provider rejects it instead of guessing a representation.
579        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/// CubeCL WebGPU tensor backend.
622///
623/// # Examples
624///
625/// ```
626/// use tenferro_gpu::webgpu::WebGpuBackend;
627///
628/// let _ctor: fn(usize) -> tenferro_tensor::Result<WebGpuBackend> = WebGpuBackend::new;
629/// ```
630///
631/// The backend is not an operation route: the operations live on
632/// [`WebGpuExecSession`], so the owner does not implement the operation traits.
633///
634/// ```compile_fail
635/// fn requires_elementwise<B: tenferro_tensor::TensorElementwise>() {}
636/// requires_elementwise::<tenferro_gpu::webgpu::WebGpuBackend>();
637/// ```
638#[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    /// Initialize a WebGPU backend for a discrete GPU ordinal.
653    ///
654    /// # Examples
655    ///
656    /// ```
657    /// use tenferro_gpu::webgpu::WebGpuBackend;
658    ///
659    /// let _ctor: fn(usize) -> tenferro_tensor::Result<WebGpuBackend> = WebGpuBackend::new;
660    /// ```
661    ///
662    /// # Errors
663    ///
664    /// Returns [`crate::Error::RuntimeState`] when no adapter/device is
665    /// available, or [`crate::Error::BackendSource`] when CubeCL initialization
666    /// fails.
667    pub fn new(device_ordinal: usize) -> crate::Result<Self> {
668        WebGpuRuntime::new(device_ordinal).map(Self::from_runtime)
669    }
670
671    /// Initialize a WebGPU backend using CubeCL's default adapter selection.
672    ///
673    /// # Examples
674    ///
675    /// ```
676    /// use tenferro_gpu::webgpu::WebGpuBackend;
677    ///
678    /// let _ctor: fn() -> tenferro_tensor::Result<WebGpuBackend> = WebGpuBackend::new_default;
679    /// ```
680    ///
681    /// # Errors
682    ///
683    /// Returns [`crate::Error::RuntimeState`] when default adapter selection
684    /// is unavailable, or [`crate::Error::BackendSource`] when initialization
685    /// fails.
686    pub fn new_default() -> crate::Result<Self> {
687        WebGpuRuntime::new_default().map(Self::from_runtime)
688    }
689
690    /// Build a WebGPU backend from an initialized runtime.
691    ///
692    /// # Examples
693    ///
694    /// ```
695    /// use tenferro_gpu::{webgpu::WebGpuBackend, webgpu::WebGpuRuntime};
696    ///
697    /// let _from_runtime: fn(WebGpuRuntime) -> WebGpuBackend = WebGpuBackend::from_runtime;
698    /// ```
699    pub fn from_runtime(runtime: WebGpuRuntime) -> Self {
700        Self { runtime }
701    }
702
703    /// Return this backend's WebGPU runtime.
704    ///
705    /// # Examples
706    ///
707    /// ```
708    /// use tenferro_gpu::{webgpu::WebGpuBackend, webgpu::WebGpuRuntime};
709    ///
710    /// let _runtime: fn(&WebGpuBackend) -> &WebGpuRuntime = WebGpuBackend::runtime;
711    /// ```
712    pub fn runtime(&self) -> &WebGpuRuntime {
713        &self.runtime
714    }
715
716    /// Return the opaque identity of this exact executable backend instance.
717    ///
718    /// Clones of a backend return the same identity. Independently initialized
719    /// backends return different identities even when they target the same
720    /// WebGPU device ordinal. This also covers Apple-backed WebGPU runtimes.
721    ///
722    /// # Examples
723    ///
724    /// ```
725    /// use tenferro_gpu::webgpu::WebGpuBackend;
726    ///
727    /// let _identity = WebGpuBackend::runtime_identity;
728    /// ```
729    pub fn runtime_identity(&self) -> WebGpuRuntimeIdentity {
730        self.runtime.runtime_identity()
731    }
732
733    /// Block until queued WebGPU work completes.
734    ///
735    /// # Examples
736    ///
737    /// ```
738    /// use tenferro_gpu::webgpu::WebGpuBackend;
739    ///
740    /// let _sync: fn(&WebGpuBackend) -> tenferro_tensor::Result<()> = WebGpuBackend::synchronize;
741    /// ```
742    ///
743    /// # Errors
744    ///
745    /// Returns [`crate::Error::BackendSource`] when queue flush or
746    /// synchronization fails, or [`crate::Error::RuntimeState`] when the
747    /// runtime has lost its device state.
748    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 {}