Skip to main content

singe_cutensor/mp/
tensor.rs

1use std::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        unsafe {
122            try_ffi!(sys::cutensorMpCreateTensorDescriptor(
123                context.as_raw(),
124                &raw mut handle,
125                rank,
126                extent_i64.as_ptr(),
127                element_stride_i64.as_ref().map_or(ptr::null(), Vec::as_ptr),
128                block_size_i64.as_ref().map_or(ptr::null(), Vec::as_ptr),
129                block_stride_i64.as_ref().map_or(ptr::null(), Vec::as_ptr),
130                ranks_per_mode_i64.as_ptr(),
131                config.partition.rank_count,
132                config.partition.ranks.map_or(ptr::null(), <[i32]>::as_ptr),
133                config.data_type.into(),
134            ))?;
135        }
136
137        if handle.is_null() {
138            return Err(Error::NullHandle);
139        }
140
141        Ok(Self {
142            handle,
143            context: context.as_context_ref(),
144            shape,
145            element_stride,
146            block_size,
147            block_stride,
148            ranks_per_mode: config.partition.ranks_per_mode.to_vec(),
149            ranks: config.partition.ranks.map(<[i32]>::to_vec),
150            rank_count: config.partition.rank_count,
151            data_type: config.data_type,
152        })
153    }
154
155    pub fn rank(&self) -> u32 {
156        self.shape.len() as u32
157    }
158
159    pub fn shape(&self) -> &[u64] {
160        &self.shape
161    }
162
163    pub fn element_stride(&self) -> Option<&[u64]> {
164        self.element_stride.as_deref()
165    }
166
167    pub fn block_size(&self) -> Option<&[u64]> {
168        self.block_size.as_deref()
169    }
170
171    pub fn block_stride(&self) -> Option<&[u64]> {
172        self.block_stride.as_deref()
173    }
174
175    pub fn ranks_per_mode(&self) -> &[u64] {
176        &self.ranks_per_mode
177    }
178
179    pub fn ranks(&self) -> Option<&[i32]> {
180        self.ranks.as_deref()
181    }
182
183    pub fn rank_count(&self) -> u32 {
184        self.rank_count
185    }
186
187    pub fn data_type(&self) -> DataType {
188        self.data_type
189    }
190
191    pub(crate) fn context(&self) -> &ContextRef<'comm> {
192        &self.context
193    }
194
195    pub const fn as_raw(&self) -> sys::cutensorMpTensorDescriptor_t {
196        self.handle
197    }
198}
199
200impl Drop for TensorDescriptor<'_> {
201    fn drop(&mut self) {
202        unsafe {
203            if let Err(err) = try_ffi!(sys::cutensorMpDestroyTensorDescriptor(self.handle)) {
204                #[cfg(debug_assertions)]
205                eprintln!("failed to destroy cutensormp tensor descriptor: {err}");
206            }
207        }
208    }
209}
210
211fn validate_rank<T>(values: &[T], rank: usize, name: &str) -> Result<()> {
212    if values.len() != rank {
213        return Err(Error::LengthMismatch {
214            name: name.into(),
215            expected: rank,
216            actual: values.len(),
217        });
218    }
219    Ok(())
220}
221
222impl TensorMode {
223    pub const fn contiguous(shape: u64) -> Self {
224        Self {
225            shape,
226            element_stride: None,
227            block_size: None,
228            block_stride: None,
229        }
230    }
231}
232
233fn collect_optional_mode_values(
234    modes: &[TensorMode],
235    name: &str,
236    value: impl Fn(&TensorMode) -> Option<u64>,
237) -> Result<Option<Vec<u64>>> {
238    let values = modes.iter().filter_map(value).collect::<Vec<_>>();
239    if values.is_empty() {
240        return Ok(None);
241    }
242    if values.len() != modes.len() {
243        return Err(Error::LengthMismatch {
244            name: name.into(),
245            expected: modes.len(),
246            actual: values.len(),
247        });
248    }
249    Ok(Some(values))
250}