Skip to main content

singe_cutensor/mg/
tensor.rs

1use std::{mem::ManuallyDrop, ptr};
2
3use singe_cuda::data_type::{DataType, DataTypeLike};
4
5use crate::{
6    error::{Error, Result},
7    mg::context::{Context, ContextRef},
8    sys, try_ffi,
9    utility::{to_i64_vec, to_u32},
10};
11
12#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
13#[non_exhaustive]
14pub enum TensorDevice {
15    Device(i32),
16    Host,
17    HostPinned,
18}
19
20impl TensorDevice {
21    pub const fn as_raw(self) -> i32 {
22        match self {
23            Self::Device(device) => device,
24            Self::Host => sys::cutensorMgHostDevice_t::CUTENSOR_MG_DEVICE_HOST as i32,
25            Self::HostPinned => sys::cutensorMgHostDevice_t::CUTENSOR_MG_DEVICE_HOST_PINNED as i32,
26        }
27    }
28}
29
30impl From<i32> for TensorDevice {
31    fn from(value: i32) -> Self {
32        Self::Device(value)
33    }
34}
35
36#[derive(Debug, Clone, Copy)]
37pub struct TensorMode {
38    pub shape: u64,
39    pub element_stride: Option<u64>,
40    pub block_size: Option<u64>,
41    pub block_stride: Option<u64>,
42}
43
44#[derive(Debug, Clone, Copy)]
45pub struct DevicePartition<'a> {
46    pub devices: &'a [TensorDevice],
47    pub count_per_mode: Option<&'a [i32]>,
48}
49
50#[derive(Debug, Clone, Copy)]
51pub struct TensorDescriptorConfig<'a> {
52    pub modes: &'a [TensorMode],
53    pub partition: DevicePartition<'a>,
54    pub data_type: DataType,
55}
56
57#[derive(Debug)]
58pub struct TensorDescriptor {
59    handle: sys::cutensorMgTensorDescriptor_t,
60    context: ContextRef,
61    shape: Vec<u64>,
62    element_stride: Option<Vec<u64>>,
63    block_size: Option<Vec<u64>>,
64    block_stride: Option<Vec<u64>>,
65    device_count: Option<Vec<i32>>,
66    devices: Vec<TensorDevice>,
67    data_type: DataType,
68}
69
70impl TensorDescriptor {
71    pub fn create_for<T: DataTypeLike>(
72        context: &Context,
73        shape: &[u64],
74        devices: &[TensorDevice],
75    ) -> Result<Self> {
76        let modes = shape
77            .iter()
78            .copied()
79            .map(TensorMode::contiguous)
80            .collect::<Vec<_>>();
81        Self::create(
82            context,
83            TensorDescriptorConfig {
84                modes: &modes,
85                partition: DevicePartition {
86                    devices,
87                    count_per_mode: None,
88                },
89                data_type: T::data_type(),
90            },
91        )
92    }
93
94    pub fn create(context: &Context, config: TensorDescriptorConfig<'_>) -> Result<Self> {
95        let rank = to_u32(config.modes.len(), "modes")?;
96        validate_optional_rank_i32(
97            config.partition.count_per_mode,
98            config.modes.len(),
99            "device_count",
100        )?;
101
102        let shape = config
103            .modes
104            .iter()
105            .map(|mode| mode.shape)
106            .collect::<Vec<_>>();
107        let element_stride =
108            collect_optional_mode_values(config.modes, "element_stride", |mode| {
109                mode.element_stride
110            })?;
111        let block_size =
112            collect_optional_mode_values(config.modes, "block_size", |mode| mode.block_size)?;
113        let block_stride =
114            collect_optional_mode_values(config.modes, "block_stride", |mode| mode.block_stride)?;
115
116        let extent_i64 = to_i64_vec(&shape, "shape")?;
117        let element_stride_i64 = element_stride
118            .as_ref()
119            .map(|values| to_i64_vec(values, "element_stride"))
120            .transpose()?;
121        let block_size_i64 = block_size
122            .as_ref()
123            .map(|values| to_i64_vec(values, "block_size"))
124            .transpose()?;
125        let block_stride_i64 = block_stride
126            .as_ref()
127            .map(|values| to_i64_vec(values, "block_stride"))
128            .transpose()?;
129        let devices_raw: Vec<_> = config
130            .partition
131            .devices
132            .iter()
133            .map(|device| device.as_raw())
134            .collect();
135        let num_devices = to_u32(devices_raw.len(), "devices")?;
136
137        let mut handle = ptr::null_mut();
138        unsafe {
139            try_ffi!(sys::cutensorMgCreateTensorDescriptor(
140                context.as_raw(),
141                &raw mut handle,
142                rank,
143                extent_i64.as_ptr(),
144                element_stride_i64.as_ref().map_or(ptr::null(), Vec::as_ptr),
145                block_size_i64.as_ref().map_or(ptr::null(), Vec::as_ptr),
146                block_stride_i64.as_ref().map_or(ptr::null(), Vec::as_ptr),
147                config
148                    .partition
149                    .count_per_mode
150                    .map_or(ptr::null(), <[i32]>::as_ptr),
151                num_devices,
152                devices_raw.as_ptr(),
153                config.data_type.into(),
154            ))?;
155        }
156
157        if handle.is_null() {
158            return Err(Error::NullHandle);
159        }
160
161        Ok(Self {
162            handle,
163            context: context.as_context_ref(),
164            shape,
165            element_stride,
166            block_size,
167            block_stride,
168            device_count: config.partition.count_per_mode.map(<[i32]>::to_vec),
169            devices: config.partition.devices.to_vec(),
170            data_type: config.data_type,
171        })
172    }
173
174    /// Wraps an existing cuTENSORMg tensor descriptor.
175    ///
176    /// # Safety
177    ///
178    /// `handle` must be a valid cuTENSORMg tensor descriptor associated with
179    /// `context`. The metadata in `config` must describe the raw descriptor.
180    /// The returned value takes ownership of `handle` and destroys it on drop.
181    pub unsafe fn from_raw(
182        handle: sys::cutensorMgTensorDescriptor_t,
183        context: &Context,
184        config: TensorDescriptorConfig<'_>,
185    ) -> Result<Self> {
186        if handle.is_null() {
187            return Err(Error::NullHandle);
188        }
189
190        let shape = config
191            .modes
192            .iter()
193            .map(|mode| mode.shape)
194            .collect::<Vec<_>>();
195        let element_stride =
196            collect_optional_mode_values(config.modes, "element_stride", |mode| {
197                mode.element_stride
198            })?;
199        let block_size =
200            collect_optional_mode_values(config.modes, "block_size", |mode| mode.block_size)?;
201        let block_stride =
202            collect_optional_mode_values(config.modes, "block_stride", |mode| mode.block_stride)?;
203
204        Ok(Self {
205            handle,
206            context: context.as_context_ref(),
207            shape,
208            element_stride,
209            block_size,
210            block_stride,
211            device_count: config.partition.count_per_mode.map(<[i32]>::to_vec),
212            devices: config.partition.devices.to_vec(),
213            data_type: config.data_type,
214        })
215    }
216
217    pub fn rank(&self) -> u32 {
218        self.shape.len() as u32
219    }
220
221    pub fn shape(&self) -> &[u64] {
222        &self.shape
223    }
224
225    pub fn element_stride(&self) -> Option<&[u64]> {
226        self.element_stride.as_deref()
227    }
228
229    pub fn block_size(&self) -> Option<&[u64]> {
230        self.block_size.as_deref()
231    }
232
233    pub fn block_stride(&self) -> Option<&[u64]> {
234        self.block_stride.as_deref()
235    }
236
237    pub fn device_count(&self) -> Option<&[i32]> {
238        self.device_count.as_deref()
239    }
240
241    pub fn devices(&self) -> &[TensorDevice] {
242        &self.devices
243    }
244
245    pub fn data_type(&self) -> DataType {
246        self.data_type
247    }
248
249    pub(crate) fn context(&self) -> &ContextRef {
250        &self.context
251    }
252
253    pub const fn as_raw(&self) -> sys::cutensorMgTensorDescriptor_t {
254        self.handle
255    }
256
257    /// Consumes this descriptor and returns the owned raw cuTENSORMg tensor descriptor.
258    ///
259    /// The caller becomes responsible for destroying the descriptor.
260    pub fn into_raw(self) -> sys::cutensorMgTensorDescriptor_t {
261        let this = ManuallyDrop::new(self);
262        this.handle
263    }
264}
265
266impl Drop for TensorDescriptor {
267    fn drop(&mut self) {
268        unsafe {
269            if let Err(err) = try_ffi!(sys::cutensorMgDestroyTensorDescriptor(self.handle)) {
270                #[cfg(debug_assertions)]
271                eprintln!("failed to destroy cutensormg tensor descriptor: {err}");
272            }
273        }
274    }
275}
276
277impl TensorMode {
278    pub const fn contiguous(shape: u64) -> Self {
279        Self {
280            shape,
281            element_stride: None,
282            block_size: None,
283            block_stride: None,
284        }
285    }
286}
287
288fn validate_optional_rank_i32(values: Option<&[i32]>, rank: usize, name: &str) -> Result<()> {
289    if let Some(values) = values
290        && values.len() != rank
291    {
292        return Err(Error::LengthMismatch {
293            name: name.into(),
294            expected: rank,
295            actual: values.len(),
296        });
297    }
298    Ok(())
299}
300
301fn collect_optional_mode_values(
302    modes: &[TensorMode],
303    name: &str,
304    value: impl Fn(&TensorMode) -> Option<u64>,
305) -> Result<Option<Vec<u64>>> {
306    let values = modes.iter().filter_map(value).collect::<Vec<_>>();
307    if values.is_empty() {
308        return Ok(None);
309    }
310    if values.len() != modes.len() {
311        return Err(Error::LengthMismatch {
312            name: name.into(),
313            expected: modes.len(),
314            actual: values.len(),
315        });
316    }
317    Ok(Some(values))
318}