Skip to main content

tenferro_gpu/cubecl/
memory.rs

1//! Host-to-device and device-to-host transfers via CubeCL-managed allocations.
2
3use 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
16/// Upload a host tensor into a CubeCL-managed GPU allocation.
17///
18/// # Examples
19///
20/// ```
21/// use tenferro_gpu::{cuda::upload_tensor, cuda::CudaRuntime};
22/// use tenferro_tensor::{Result, Tensor};
23///
24/// let _upload: fn(&CudaRuntime, &Tensor) -> Result<Tensor> = upload_tensor;
25/// ```
26///
27/// # Errors
28///
29/// Returns [`crate::Error::RuntimeState`] when the source is backend-resident
30/// or belongs to another placement, [`crate::Error::Unsupported`] for a dtype
31/// unavailable in CubeCL, or [`crate::Error::BackendSource`] on allocation.
32pub 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        // A caller-owned payload has no GPU implementation for this operation.
54        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
61/// The typed tensor behind `tensor`, or a typed refusal.
62///
63/// Callers reach this from a match on the tensor's dtype, so `None` means the tag
64/// table and the runtime dtype disagree rather than a caller mistake.
65fn 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
74/// Download a CubeCL-managed GPU tensor back to host memory.
75///
76/// # Examples
77///
78/// ```
79/// use tenferro_gpu::{cuda::download_tensor, cuda::CudaRuntime};
80/// use tenferro_tensor::{Result, Tensor};
81///
82/// let _download: fn(&CudaRuntime, &Tensor) -> Result<Tensor> = download_tensor;
83/// ```
84///
85/// # Errors
86///
87/// Returns [`crate::Error::RuntimeState`] for a host-backed or foreign tensor,
88/// [`crate::Error::BackendSource`] when synchronization/readback fails, or a
89/// typed validation error when device data cannot be decoded.
90pub 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        // A caller-owned payload has no GPU implementation for this operation.
112        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    // The host buffer is borrowed, so CubeCL stages exactly one copy of it: the
138    // device write runs after this call returns (#2009).
139    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
160/// Read a compact tensor of at most `PINNED_SCALAR_BYTES` through the runtime's
161/// pinned staging slot.
162fn 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    // Small-payload fast path. A Krylov loop reads back one reduction result
197    // per iteration. CubeCL's `read_one` already stages through pinned memory,
198    // so this is not a pageable-to-pinned conversion; what it avoids is
199    // CubeCL's staging-reservation machinery and one of two barriers. The
200    // general path below synchronizes the stream and then lets `read_one` wait
201    // on its own fence after the copy, while this issues the copy and waits
202    // once. Neither path synchronizes the device. The gate is a byte length, so
203    // a short vector takes it too, not only a scalar.
204    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        // A handle that does not span exactly the tensor's elements keeps the
223        // whole-allocation read and its element-count validation below.
224        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
238/// Download a whole allocation of `len` elements into a freshly owned `Vec<T>`.
239///
240/// The device copies straight into the vector the host tensor will own, so the
241/// payload crosses host memory once instead of being staged by CubeCL and then
242/// copied again (#2009). The vector is allocated with its final element type,
243/// so its alignment is `T`'s, and it is only exposed after the copy completed.
244pub(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    // SAFETY: the spare capacity of `data` holds `len` elements, i.e.
253    // `byte_len` bytes, and `data` outlives the call. On error the device may
254    // still write into it, so it is leaked below rather than freed.
255    let copied = unsafe { rt.download_into_host(handle, data.as_mut_ptr().cast(), byte_len, op) };
256    match copied {
257        Ok(()) => {
258            // SAFETY: the completed copy initialized all `len` elements, and
259            // `T: CubeElement` is `Pod`, so any byte pattern is a valid `T`.
260            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    // The converted bytes are owned: hand them to CubeCL without a staging copy.
291    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        // A caller-owned payload has no GPU implementation for this operation.
377        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}