1use cubecl::client::ComputeClient;
4use cubecl::prelude::CubeElement;
5use cubecl_cuda::CudaRuntime as CubeclCudaRuntime;
6use num_complex::{Complex32, Complex64};
7use tenferro_tensor::DType;
8
9use super::dispatch;
10use crate::cubecl::runtime::{CudaRuntime, PINNED_SCALAR_BYTES};
11use crate::types::{
12 CubeclBuffer, DeviceId, DeviceKind, GpuBackendKind, MemoryKind, Placement, StorageBuffer,
13 Tensor, TensorScalar, TypedTensor,
14};
15
16pub fn upload_tensor(rt: &CudaRuntime, tensor: &Tensor) -> crate::Result<Tensor> {
33 let client = rt.client();
34 match tensor.dtype() {
35 DType::F64 => upload_typed::<f64>(rt, client, gpu_typed::<f64>("upload_tensor", tensor)?)
36 .map(Tensor::from_typed::<f64>),
37 DType::F32 => upload_typed::<f32>(rt, client, gpu_typed::<f32>("upload_tensor", tensor)?)
38 .map(Tensor::from_typed::<f32>),
39 DType::I32 => upload_typed::<i32>(rt, client, gpu_typed::<i32>("upload_tensor", tensor)?)
40 .map(Tensor::from_typed::<i32>),
41 DType::I64 => upload_typed::<i64>(rt, client, gpu_typed::<i64>("upload_tensor", tensor)?)
42 .map(Tensor::from_typed::<i64>),
43 DType::Bool => upload_bool(rt, client, gpu_typed::<bool>("upload_tensor", tensor)?)
44 .map(Tensor::from_typed::<bool>),
45 DType::C64 => {
46 upload_typed::<Complex64>(rt, client, gpu_typed::<Complex64>("upload_tensor", tensor)?)
47 .map(Tensor::from_typed::<tenferro_tensor::Complex64>)
48 }
49 DType::C32 => {
50 upload_typed::<Complex32>(rt, client, gpu_typed::<Complex32>("upload_tensor", tensor)?)
51 .map(Tensor::from_typed::<tenferro_tensor::Complex32>)
52 }
53 DType::External(_) => Err(crate::Error::unsupported(
55 "upload_tensor",
56 "an externally defined payload is not supported by this GPU operation",
57 )),
58 }
59}
60
61fn gpu_typed<'a, T: TensorScalar>(
66 op: &'static str,
67 tensor: &'a Tensor,
68) -> crate::Result<&'a TypedTensor<T>> {
69 tensor.as_typed::<T>().ok_or_else(|| {
70 crate::Error::unsupported(op, "the GPU memory path requires a preset scalar")
71 })
72}
73
74pub fn download_tensor(rt: &CudaRuntime, tensor: &Tensor) -> crate::Result<Tensor> {
91 ensure_tensor_resident_on_runtime(rt, tensor, "download")?;
92 match tensor.dtype() {
93 DType::F64 => download_typed::<f64>(rt, gpu_typed::<f64>("download_tensor", tensor)?)
94 .map(Tensor::from_typed::<f64>),
95 DType::F32 => download_typed::<f32>(rt, gpu_typed::<f32>("download_tensor", tensor)?)
96 .map(Tensor::from_typed::<f32>),
97 DType::I32 => download_typed::<i32>(rt, gpu_typed::<i32>("download_tensor", tensor)?)
98 .map(Tensor::from_typed::<i32>),
99 DType::I64 => download_typed::<i64>(rt, gpu_typed::<i64>("download_tensor", tensor)?)
100 .map(Tensor::from_typed::<i64>),
101 DType::Bool => download_bool(rt, gpu_typed::<bool>("download_tensor", tensor)?)
102 .map(Tensor::from_typed::<bool>),
103 DType::C64 => {
104 download_typed::<Complex64>(rt, gpu_typed::<Complex64>("download_tensor", tensor)?)
105 .map(Tensor::from_typed::<tenferro_tensor::Complex64>)
106 }
107 DType::C32 => {
108 download_typed::<Complex32>(rt, gpu_typed::<Complex32>("download_tensor", tensor)?)
109 .map(Tensor::from_typed::<tenferro_tensor::Complex32>)
110 }
111 DType::External(_) => Err(crate::Error::unsupported(
113 "download_tensor",
114 "an externally defined payload is not supported by this GPU operation",
115 )),
116 }
117}
118
119fn upload_typed<T: CubeElement + TensorScalar + Clone + Send + Sync + 'static>(
120 rt: &CudaRuntime,
121 client: &ComputeClient<CubeclCudaRuntime>,
122 typed: &TypedTensor<T>,
123) -> crate::Result<TypedTensor<T>> {
124 let host_data = match typed.buffer() {
125 StorageBuffer::Host(data) => data,
126 StorageBuffer::Backend(buffer) => {
127 return Err(crate::Error::runtime_state(
128 "upload",
129 format!(
130 "expected host buffer, got `{}` backend buffer",
131 buffer.backend_family()
132 ),
133 ));
134 }
135 };
136
137 let handle = client.create_from_slice(T::as_bytes(host_data));
140 let byte_len = T::as_bytes(host_data).len();
141 TypedTensor::from_buffer_col_major(
142 typed.shape().to_vec(),
143 StorageBuffer::Backend(Box::new(CubeclBuffer::new(
144 handle,
145 byte_len,
146 rt.device_ordinal(),
147 rt.allocation_domain_id(),
148 ))),
149 Placement {
150 memory_kind: MemoryKind::Device,
151 device: Some(DeviceId {
152 kind: DeviceKind::Gpu(GpuBackendKind::Cuda),
153 ordinal: rt.device_ordinal(),
154 }),
155 cpu_affinity: None,
156 },
157 )
158}
159
160fn download_scalar<T: CubeElement + TensorScalar + Clone + 'static>(
163 rt: &CudaRuntime,
164 typed: &TypedTensor<T>,
165 byte_len: usize,
166) -> crate::Result<Vec<T>> {
167 let ptr = super::gemm::typed_device_ptr(rt, typed, "download")?;
168 let retained = dispatch::cubecl_buffer(typed, "download")?.handle().clone();
169 let mut bytes = vec![0_u8; byte_len];
170 rt.download_scalar_bytes(ptr as u64, &mut bytes, "download", retained)?;
171 Ok(T::from_bytes(&bytes).to_vec())
172}
173
174fn download_typed<T: CubeElement + TensorScalar + Clone + 'static>(
175 rt: &CudaRuntime,
176 typed: &TypedTensor<T>,
177) -> crate::Result<TypedTensor<T>> {
178 let handle = match typed.buffer() {
179 StorageBuffer::Host(_) => {
180 return Err(crate::Error::runtime_state(
181 "download",
182 "expected CubeCL buffer",
183 ));
184 }
185 StorageBuffer::Backend(buffer) => cubecl_handle_from_backend(buffer.as_ref(), "download")?,
186 };
187
188 if typed.n_elements() == 0 {
189 return TypedTensor::from_buffer_col_major(
190 typed.shape().to_vec(),
191 StorageBuffer::Host(Vec::new()),
192 Placement::default(),
193 );
194 }
195
196 let byte_len = typed
205 .n_elements()
206 .checked_mul(size_of::<T>())
207 .ok_or_else(|| {
208 crate::Error::invalid_argument("download", "shape", "byte length overflows")
209 })?;
210 if byte_len <= PINNED_SCALAR_BYTES {
211 let data = download_scalar(rt, typed, byte_len)?;
212 return TypedTensor::from_buffer_col_major(
213 typed.shape().to_vec(),
214 StorageBuffer::Host(data),
215 Placement::default(),
216 );
217 }
218
219 let data = if handle.size_in_used() == byte_len as u64 {
220 download_owned_vec::<T>(rt, handle, typed.n_elements(), byte_len, "download")?
221 } else {
222 rt.synchronize()?;
225 let bytes = rt
226 .client()
227 .read_one(handle)
228 .map_err(|err| crate::Error::backend_source("download", err))?;
229 T::from_bytes(&bytes).to_vec()
230 };
231 TypedTensor::from_buffer_col_major(
232 typed.shape().to_vec(),
233 StorageBuffer::Host(data),
234 Placement::default(),
235 )
236}
237
238pub(super) fn download_owned_vec<T: CubeElement>(
245 rt: &CudaRuntime,
246 handle: cubecl_runtime::server::Handle,
247 len: usize,
248 byte_len: usize,
249 op: &'static str,
250) -> crate::Result<Vec<T>> {
251 let mut data = Vec::<T>::with_capacity(len);
252 let copied = unsafe { rt.download_into_host(handle, data.as_mut_ptr().cast(), byte_len, op) };
256 match copied {
257 Ok(()) => {
258 unsafe { data.set_len(len) };
261 Ok(data)
262 }
263 Err(err) => {
264 std::mem::forget(data);
265 Err(err)
266 }
267 }
268}
269
270fn upload_bool(
271 rt: &CudaRuntime,
272 client: &ComputeClient<CubeclCudaRuntime>,
273 typed: &TypedTensor<bool>,
274) -> crate::Result<TypedTensor<bool>> {
275 let host_data = match typed.buffer() {
276 StorageBuffer::Host(data) => data,
277 StorageBuffer::Backend(buffer) => {
278 return Err(crate::Error::runtime_state(
279 "upload",
280 format!(
281 "expected host buffer, got `{}` backend buffer",
282 buffer.backend_family()
283 ),
284 ));
285 }
286 };
287
288 let bytes: Vec<u8> = host_data.iter().map(|&value| u8::from(value)).collect();
289 let byte_len = bytes.len();
290 let handle = client.create(cubecl::bytes::Bytes::from_elems(bytes));
292 TypedTensor::from_buffer_col_major(
293 typed.shape().to_vec(),
294 StorageBuffer::Backend(Box::new(CubeclBuffer::new(
295 handle,
296 byte_len,
297 rt.device_ordinal(),
298 rt.allocation_domain_id(),
299 ))),
300 Placement {
301 memory_kind: MemoryKind::Device,
302 device: Some(DeviceId {
303 kind: DeviceKind::Gpu(GpuBackendKind::Cuda),
304 ordinal: rt.device_ordinal(),
305 }),
306 cpu_affinity: None,
307 },
308 )
309}
310
311fn download_bool(rt: &CudaRuntime, typed: &TypedTensor<bool>) -> crate::Result<TypedTensor<bool>> {
312 let handle = match typed.buffer() {
313 StorageBuffer::Host(_) => {
314 return Err(crate::Error::runtime_state(
315 "download",
316 "expected CubeCL buffer",
317 ));
318 }
319 StorageBuffer::Backend(buffer) => cubecl_handle_from_backend(buffer.as_ref(), "download")?,
320 };
321
322 if typed.n_elements() == 0 {
323 return TypedTensor::from_vec_col_major(typed.shape().to_vec(), Vec::new());
324 }
325
326 rt.synchronize()?;
327 let bytes = rt
328 .client()
329 .read_one(handle)
330 .map_err(|err| crate::Error::backend_source("download", err))?;
331 let data = bytes.iter().map(|&byte| byte != 0).collect();
332 TypedTensor::from_vec_col_major(typed.shape().to_vec(), data)
333}
334
335fn ensure_tensor_resident_on_runtime(
336 rt: &CudaRuntime,
337 tensor: &Tensor,
338 op: &'static str,
339) -> crate::Result<()> {
340 match tensor.dtype() {
341 DType::F64 => dispatch::ensure_resident_on_runtime(
342 rt,
343 gpu_typed::<f64>("ensure_tensor_resident_on_runtime", tensor)?,
344 op,
345 ),
346 DType::F32 => dispatch::ensure_resident_on_runtime(
347 rt,
348 gpu_typed::<f32>("ensure_tensor_resident_on_runtime", tensor)?,
349 op,
350 ),
351 DType::I32 => dispatch::ensure_resident_on_runtime(
352 rt,
353 gpu_typed::<i32>("ensure_tensor_resident_on_runtime", tensor)?,
354 op,
355 ),
356 DType::I64 => dispatch::ensure_resident_on_runtime(
357 rt,
358 gpu_typed::<i64>("ensure_tensor_resident_on_runtime", tensor)?,
359 op,
360 ),
361 DType::Bool => dispatch::ensure_resident_on_runtime(
362 rt,
363 gpu_typed::<bool>("ensure_tensor_resident_on_runtime", tensor)?,
364 op,
365 ),
366 DType::C64 => dispatch::ensure_resident_on_runtime(
367 rt,
368 gpu_typed::<Complex64>("ensure_tensor_resident_on_runtime", tensor)?,
369 op,
370 ),
371 DType::C32 => dispatch::ensure_resident_on_runtime(
372 rt,
373 gpu_typed::<Complex32>("ensure_tensor_resident_on_runtime", tensor)?,
374 op,
375 ),
376 DType::External(_) => Err(crate::Error::unsupported(
378 "ensure_tensor_resident_on_runtime",
379 "an externally defined payload is not supported by this GPU operation",
380 )),
381 }
382}
383
384fn cubecl_handle_from_backend<T: 'static>(
385 buffer: &dyn crate::BackendStorage<T>,
386 op: &'static str,
387) -> crate::Result<cubecl_runtime::server::Handle> {
388 buffer
389 .as_any()
390 .downcast_ref::<CubeclBuffer>()
391 .map(|buffer| buffer.handle().clone())
392 .ok_or_else(|| {
393 crate::Error::runtime_state(
394 op,
395 format!(
396 "expected CubeCL buffer, got `{}` backend buffer",
397 buffer.backend_family()
398 ),
399 )
400 })
401}