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
13fn 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
26pub 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 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
73pub 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 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}