Skip to main content

tenferro_gpu/webgpu/
memory.rs

1use cubecl::client::ComputeClient;
2use cubecl::prelude::CubeElement;
3use cubecl_wgpu::WgpuRuntime;
4use num_complex::{Complex32, Complex64};
5use tenferro_tensor::DType;
6
7use super::{
8    ensure_resident_on_runtime, prepared_webgpu_tensor, typed_from_webgpu, WebGpuBuffer,
9    WebGpuRuntime,
10};
11use crate::{Tensor, TypedTensor};
12
13/// The typed tensor behind `tensor`, or a typed refusal.
14///
15/// Callers reach this from a match on the tensor's dtype, so `None` means the tag
16/// table and the runtime dtype disagree rather than a caller mistake.
17fn webgpu_typed<'a, T: tenferro_tensor::TensorScalar>(
18    op: &'static str,
19    tensor: &'a Tensor,
20) -> crate::Result<&'a TypedTensor<T>> {
21    tensor.as_typed::<T>().ok_or_else(|| {
22        crate::Error::unsupported(op, "the WebGPU memory path requires a preset scalar")
23    })
24}
25
26/// Upload a host tensor into a CubeCL-managed WebGPU allocation.
27///
28/// # Examples
29///
30/// ```
31/// use tenferro_gpu::{webgpu::upload_webgpu_tensor, webgpu::WebGpuRuntime};
32/// use tenferro_tensor::{Result, Tensor};
33///
34/// let _upload: fn(&WebGpuRuntime, &Tensor) -> Result<Tensor> = upload_webgpu_tensor;
35/// ```
36///
37/// # Errors
38///
39/// Returns [`crate::Error::RuntimeState`] when the source buffer is backend
40/// resident or belongs to another placement, [`crate::Error::Unsupported`] for
41/// a dtype unavailable in WebGPU, or [`crate::Error::BackendSource`] on
42/// allocation.
43pub fn upload_webgpu_tensor(rt: &WebGpuRuntime, tensor: &Tensor) -> crate::Result<Tensor> {
44    match tensor.dtype() {
45        DType::F64 => upload_typed::<f64>(rt, webgpu_typed::<f64>("upload_webgpu_tensor", tensor)?)
46            .map(Tensor::from_typed::<f64>),
47        DType::F32 => upload_typed::<f32>(rt, webgpu_typed::<f32>("upload_webgpu_tensor", tensor)?)
48            .map(Tensor::from_typed::<f32>),
49        DType::I32 => upload_typed::<i32>(rt, webgpu_typed::<i32>("upload_webgpu_tensor", tensor)?)
50            .map(Tensor::from_typed::<i32>),
51        DType::I64 => upload_typed::<i64>(rt, webgpu_typed::<i64>("upload_webgpu_tensor", tensor)?)
52            .map(Tensor::from_typed::<i64>),
53        DType::Bool => upload_bool(rt, webgpu_typed::<bool>("upload_webgpu_tensor", tensor)?)
54            .map(Tensor::from_typed::<bool>),
55        DType::C64 => upload_typed::<Complex64>(
56            rt,
57            webgpu_typed::<Complex64>("upload_webgpu_tensor", tensor)?,
58        )
59        .map(Tensor::from_typed::<num_complex::Complex64>),
60        DType::C32 => upload_typed::<Complex32>(
61            rt,
62            webgpu_typed::<Complex32>("upload_webgpu_tensor", tensor)?,
63        )
64        .map(Tensor::from_typed::<num_complex::Complex32>),
65        // A caller-owned payload has no GPU implementation for this operation.
66        DType::External(_) => Err(crate::Error::unsupported(
67            "upload_webgpu_tensor",
68            "an externally defined payload is not supported by this GPU operation",
69        )),
70    }
71}
72
73/// Download a CubeCL-managed WebGPU tensor back to host memory.
74///
75/// # Examples
76///
77/// ```
78/// use tenferro_gpu::{webgpu::download_webgpu_tensor, webgpu::WebGpuRuntime};
79/// use tenferro_tensor::{Result, Tensor};
80///
81/// let _download: fn(&WebGpuRuntime, &Tensor) -> Result<Tensor> = download_webgpu_tensor;
82/// ```
83///
84/// # Errors
85///
86/// Returns [`crate::Error::RuntimeState`] for missing or foreign device state,
87/// [`crate::Error::BackendSource`] when queue synchronization/readback fails,
88/// or a typed validation error when bytes do not match the tensor shape.
89pub fn download_webgpu_tensor(rt: &WebGpuRuntime, tensor: &Tensor) -> crate::Result<Tensor> {
90    let client = rt.client();
91    match tensor.dtype() {
92        DType::F64 => download_typed::<f64>(
93            rt,
94            client,
95            webgpu_typed::<f64>("download_webgpu_tensor", tensor)?,
96        )
97        .map(Tensor::from_typed::<f64>),
98        DType::F32 => download_typed::<f32>(
99            rt,
100            client,
101            webgpu_typed::<f32>("download_webgpu_tensor", tensor)?,
102        )
103        .map(Tensor::from_typed::<f32>),
104        DType::I32 => download_typed::<i32>(
105            rt,
106            client,
107            webgpu_typed::<i32>("download_webgpu_tensor", tensor)?,
108        )
109        .map(Tensor::from_typed::<i32>),
110        DType::I64 => download_typed::<i64>(
111            rt,
112            client,
113            webgpu_typed::<i64>("download_webgpu_tensor", tensor)?,
114        )
115        .map(Tensor::from_typed::<i64>),
116        DType::Bool => download_bool(
117            rt,
118            client,
119            webgpu_typed::<bool>("download_webgpu_tensor", tensor)?,
120        )
121        .map(Tensor::from_typed::<bool>),
122        DType::C64 => download_typed::<Complex64>(
123            rt,
124            client,
125            webgpu_typed::<Complex64>("download_webgpu_tensor", tensor)?,
126        )
127        .map(Tensor::from_typed::<num_complex::Complex64>),
128        DType::C32 => download_typed::<Complex32>(
129            rt,
130            client,
131            webgpu_typed::<Complex32>("download_webgpu_tensor", tensor)?,
132        )
133        .map(Tensor::from_typed::<num_complex::Complex32>),
134        // A caller-owned payload has no GPU implementation for this operation.
135        DType::External(_) => Err(crate::Error::unsupported(
136            "download_webgpu_tensor",
137            "an externally defined payload is not supported by this GPU operation",
138        )),
139    }
140}
141
142pub(super) fn upload_typed<T: CubeElement + crate::TensorScalar + Clone + Send + Sync + 'static>(
143    rt: &WebGpuRuntime,
144    typed: &TypedTensor<T>,
145) -> crate::Result<TypedTensor<T>> {
146    let host_data = typed.host_data().map_err(|error| {
147        crate::Error::runtime_state("webgpu_upload", format!("expected host buffer: {error}"))
148    })?;
149
150    let byte_len = T::as_bytes(host_data).len();
151    let handle = rt.client().create_from_slice(T::as_bytes(host_data));
152    let buffer = WebGpuBuffer::new_for_runtime(rt, handle, byte_len, "webgpu_upload")?;
153    let tensor = typed_from_webgpu(typed.shape().to_vec(), buffer, rt)?;
154    rt.record_upload(byte_len);
155    Ok(tensor)
156}
157
158pub(super) fn download_typed<T: CubeElement + crate::TensorScalar + Clone + 'static>(
159    rt: &WebGpuRuntime,
160    client: &ComputeClient<WgpuRuntime>,
161    typed: &TypedTensor<T>,
162) -> crate::Result<TypedTensor<T>> {
163    ensure_resident_on_runtime(rt, typed, "webgpu_download")?;
164    let handle = prepared_webgpu_tensor(typed, "webgpu_download")?.handle;
165
166    if typed.n_elements() == 0 {
167        return TypedTensor::from_vec_col_major(typed.shape().to_vec(), Vec::new());
168    }
169
170    let bytes = client
171        .read_one(handle)
172        .map_err(|err| crate::Error::backend_source("webgpu_download", err))?;
173    let data = T::from_bytes(&bytes).to_vec();
174    rt.record_download(bytes.len());
175    TypedTensor::from_vec_col_major(typed.shape().to_vec(), data)
176}
177
178fn upload_bool(rt: &WebGpuRuntime, typed: &TypedTensor<bool>) -> crate::Result<TypedTensor<bool>> {
179    let host_data = typed.host_data().map_err(|error| {
180        crate::Error::runtime_state("webgpu_upload", format!("expected host buffer: {error}"))
181    })?;
182
183    let bytes: Vec<u8> = host_data.iter().map(|&value| u8::from(value)).collect();
184    let handle = rt.client().create_from_slice(&bytes);
185    let buffer = WebGpuBuffer::new_for_runtime(rt, handle, bytes.len(), "webgpu_upload")?;
186    rt.record_upload(bytes.len());
187    TypedTensor::from_backend_allocation(
188        typed.shape().to_vec(),
189        Box::new(buffer),
190        super::webgpu_placement(rt),
191    )
192}
193
194fn download_bool(
195    rt: &WebGpuRuntime,
196    client: &ComputeClient<WgpuRuntime>,
197    typed: &TypedTensor<bool>,
198) -> crate::Result<TypedTensor<bool>> {
199    ensure_resident_on_runtime(rt, typed, "webgpu_download")?;
200    let handle = prepared_webgpu_tensor(typed, "webgpu_download")?.handle;
201
202    if typed.n_elements() == 0 {
203        return TypedTensor::from_vec_col_major(typed.shape().to_vec(), Vec::new());
204    }
205
206    let bytes = client
207        .read_one(handle)
208        .map_err(|err| crate::Error::backend_source("webgpu_download", err))?;
209    let data = bytes.iter().map(|&byte| byte != 0).collect();
210    rt.record_download(bytes.len());
211    TypedTensor::from_vec_col_major(typed.shape().to_vec(), data)
212}