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 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 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}