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