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