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}