Skip to main content

singe_cutensor/mp/
tensor.rs

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