Skip to main content

singe_cutensor/mg/
tensor.rs

1use std::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)]
13pub enum TensorDevice {
14    Device(i32),
15    Host,
16    HostPinned,
17}
18
19impl TensorDevice {
20    pub const fn as_raw(self) -> i32 {
21        match self {
22            Self::Device(device) => device,
23            Self::Host => sys::cutensorMgHostDevice_t::CUTENSOR_MG_DEVICE_HOST as i32,
24            Self::HostPinned => sys::cutensorMgHostDevice_t::CUTENSOR_MG_DEVICE_HOST_PINNED as i32,
25        }
26    }
27}
28
29impl From<i32> for TensorDevice {
30    fn from(value: i32) -> Self {
31        Self::Device(value)
32    }
33}
34
35#[derive(Debug, Clone, Copy)]
36pub struct TensorMode {
37    pub shape: u64,
38    pub element_stride: Option<u64>,
39    pub block_size: Option<u64>,
40    pub block_stride: Option<u64>,
41}
42
43#[derive(Debug, Clone, Copy)]
44pub struct DevicePartition<'a> {
45    pub devices: &'a [TensorDevice],
46    pub count_per_mode: Option<&'a [i32]>,
47}
48
49#[derive(Debug, Clone, Copy)]
50pub struct TensorDescriptorConfig<'a> {
51    pub modes: &'a [TensorMode],
52    pub partition: DevicePartition<'a>,
53    pub data_type: DataType,
54}
55
56#[derive(Debug)]
57pub struct TensorDescriptor {
58    handle: sys::cutensorMgTensorDescriptor_t,
59    context: ContextRef,
60    shape: Vec<u64>,
61    element_stride: Option<Vec<u64>>,
62    block_size: Option<Vec<u64>>,
63    block_stride: Option<Vec<u64>>,
64    device_count: Option<Vec<i32>>,
65    devices: Vec<TensorDevice>,
66    data_type: DataType,
67}
68
69impl TensorDescriptor {
70    pub fn create_for<T: DataTypeLike>(
71        context: &Context,
72        shape: &[u64],
73        devices: &[TensorDevice],
74    ) -> Result<Self> {
75        let modes = shape
76            .iter()
77            .copied()
78            .map(TensorMode::contiguous)
79            .collect::<Vec<_>>();
80        Self::create(
81            context,
82            TensorDescriptorConfig {
83                modes: &modes,
84                partition: DevicePartition {
85                    devices,
86                    count_per_mode: None,
87                },
88                data_type: T::data_type(),
89            },
90        )
91    }
92
93    pub fn create(context: &Context, config: TensorDescriptorConfig<'_>) -> Result<Self> {
94        let rank = to_u32(config.modes.len(), "modes")?;
95        validate_optional_rank_i32(
96            config.partition.count_per_mode,
97            config.modes.len(),
98            "device_count",
99        )?;
100
101        let shape = config
102            .modes
103            .iter()
104            .map(|mode| mode.shape)
105            .collect::<Vec<_>>();
106        let element_stride =
107            collect_optional_mode_values(config.modes, "element_stride", |mode| {
108                mode.element_stride
109            })?;
110        let block_size =
111            collect_optional_mode_values(config.modes, "block_size", |mode| mode.block_size)?;
112        let block_stride =
113            collect_optional_mode_values(config.modes, "block_stride", |mode| mode.block_stride)?;
114
115        let extent_i64 = to_i64_vec(&shape, "shape")?;
116        let element_stride_i64 = element_stride
117            .as_ref()
118            .map(|values| to_i64_vec(values, "element_stride"))
119            .transpose()?;
120        let block_size_i64 = block_size
121            .as_ref()
122            .map(|values| to_i64_vec(values, "block_size"))
123            .transpose()?;
124        let block_stride_i64 = block_stride
125            .as_ref()
126            .map(|values| to_i64_vec(values, "block_stride"))
127            .transpose()?;
128        let devices_raw: Vec<_> = config
129            .partition
130            .devices
131            .iter()
132            .map(|device| device.as_raw())
133            .collect();
134        let num_devices = to_u32(devices_raw.len(), "devices")?;
135
136        let mut handle = ptr::null_mut();
137        unsafe {
138            try_ffi!(sys::cutensorMgCreateTensorDescriptor(
139                context.as_raw(),
140                &raw mut handle,
141                rank,
142                extent_i64.as_ptr(),
143                element_stride_i64.as_ref().map_or(ptr::null(), Vec::as_ptr),
144                block_size_i64.as_ref().map_or(ptr::null(), Vec::as_ptr),
145                block_stride_i64.as_ref().map_or(ptr::null(), Vec::as_ptr),
146                config
147                    .partition
148                    .count_per_mode
149                    .map_or(ptr::null(), <[i32]>::as_ptr),
150                num_devices,
151                devices_raw.as_ptr(),
152                config.data_type.into(),
153            ))?;
154        }
155
156        if handle.is_null() {
157            return Err(Error::NullHandle);
158        }
159
160        Ok(Self {
161            handle,
162            context: context.as_context_ref(),
163            shape,
164            element_stride,
165            block_size,
166            block_stride,
167            device_count: config.partition.count_per_mode.map(<[i32]>::to_vec),
168            devices: config.partition.devices.to_vec(),
169            data_type: config.data_type,
170        })
171    }
172
173    pub fn rank(&self) -> u32 {
174        self.shape.len() as u32
175    }
176
177    pub fn shape(&self) -> &[u64] {
178        &self.shape
179    }
180
181    pub fn element_stride(&self) -> Option<&[u64]> {
182        self.element_stride.as_deref()
183    }
184
185    pub fn block_size(&self) -> Option<&[u64]> {
186        self.block_size.as_deref()
187    }
188
189    pub fn block_stride(&self) -> Option<&[u64]> {
190        self.block_stride.as_deref()
191    }
192
193    pub fn device_count(&self) -> Option<&[i32]> {
194        self.device_count.as_deref()
195    }
196
197    pub fn devices(&self) -> &[TensorDevice] {
198        &self.devices
199    }
200
201    pub fn data_type(&self) -> DataType {
202        self.data_type
203    }
204
205    pub(crate) fn context(&self) -> &ContextRef {
206        &self.context
207    }
208
209    pub const fn as_raw(&self) -> sys::cutensorMgTensorDescriptor_t {
210        self.handle
211    }
212}
213
214impl Drop for TensorDescriptor {
215    fn drop(&mut self) {
216        unsafe {
217            if let Err(err) = try_ffi!(sys::cutensorMgDestroyTensorDescriptor(self.handle)) {
218                #[cfg(debug_assertions)]
219                eprintln!("failed to destroy cutensormg tensor descriptor: {err}");
220            }
221        }
222    }
223}
224
225impl TensorMode {
226    pub const fn contiguous(shape: u64) -> Self {
227        Self {
228            shape,
229            element_stride: None,
230            block_size: None,
231            block_stride: None,
232        }
233    }
234}
235
236fn validate_optional_rank_i32(values: Option<&[i32]>, rank: usize, name: &str) -> Result<()> {
237    if let Some(values) = values
238        && values.len() != rank
239    {
240        return Err(Error::LengthMismatch {
241            name: name.into(),
242            expected: rank,
243            actual: values.len(),
244        });
245    }
246    Ok(())
247}
248
249fn collect_optional_mode_values(
250    modes: &[TensorMode],
251    name: &str,
252    value: impl Fn(&TensorMode) -> Option<u64>,
253) -> Result<Option<Vec<u64>>> {
254    let values = modes.iter().filter_map(value).collect::<Vec<_>>();
255    if values.is_empty() {
256        return Ok(None);
257    }
258    if values.len() != modes.len() {
259        return Err(Error::LengthMismatch {
260            name: name.into(),
261            expected: modes.len(),
262            actual: values.len(),
263        });
264    }
265    Ok(Some(values))
266}