1use std::sync::{Arc, RwLock, TryLockError};
2
3use dlpk::sys::{DLDevice, DLPackVersion, DLDataType};
4use dlpk::{DLDataTypeCode, DLPackPointerCast, DLPackTensor, GetDLPackDataType, ReadOnly};
5
6use crate::errors::Error;
7use crate::c_api::mts_data_movement_t;
8
9use super::{Array, MtsArray};
10
11impl<T, D> From<ndarray::Array<T, D>> for MtsArray where
12 T: 'static + Clone + Send + Default + Sync + GetDLPackDataType + DLPackPointerCast,
13 D: ndarray::Dimension
14{
15 fn from(value: ndarray::Array<T, D>) -> Self {
16 let array = Arc::new(RwLock::new(value.into_dyn()));
17 let boxed: Box<dyn Array> = Box::new(array);
18 return MtsArray::from(boxed);
19 }
20}
21
22impl<T> Array for Arc<RwLock<ndarray::ArrayD<T>>>
23where
24 T: 'static + Send + Sync + Clone + Default + GetDLPackDataType + DLPackPointerCast,
25{
26 fn as_any(&self) -> &dyn std::any::Any {
27 self
28 }
29
30 fn as_any_mut(&mut self) -> &mut dyn std::any::Any {
31 self
32 }
33
34 fn create(&self, shape: &[usize], fill_value: MtsArray) -> Box<dyn Array> {
35 let cpu_device = DLDevice::cpu();
36 let max_version = DLPackVersion::current();
37 let fill_value_dlpack = fill_value.as_dlpack(cpu_device, None, max_version).expect("failed to extract fill_value as DLPack");
38
39 assert_eq!(fill_value_dlpack.shape(), &[], "fill_value must be a single scalar");
41 assert_eq!(fill_value_dlpack.device(), cpu_device, "fill_value must be on CPU");
42
43 let fill_value_ptr = fill_value_dlpack.data_ptr::<T>().expect("dtype mismatch between array and fill_value");
44 let fill_value_scalar = unsafe { std::ptr::read(fill_value_ptr) };
45
46 let array = ndarray::Array::from_elem(shape, fill_value_scalar);
47 return Box::new(Arc::new(RwLock::new(array)));
48 }
49
50 fn copy(&self, device: DLDevice) -> Box<dyn Array> {
51 assert_eq!(device, DLDevice::cpu(), "Rust ndarray data can only be copied to CPU device");
52 let data = self.read().expect("array lock is poisoned");
53 return Box::new(Arc::new(RwLock::new(data.clone())));
54 }
55
56 fn shape(&self) -> Vec<usize> {
57 match self.try_read() {
58 Ok(lock) => lock.shape().to_vec(),
59 Err(TryLockError::Poisoned(_)) => panic!("array lock is poisoned"),
60 Err(TryLockError::WouldBlock) => panic!("array is already locked"),
61 }
62 }
63
64 fn reshape(&mut self, shape: &[usize]) {
65 let mut lock = match self.try_write() {
66 Ok(lock) => lock,
67 Err(TryLockError::Poisoned(_)) => panic!("array lock is poisoned"),
68 Err(TryLockError::WouldBlock) => panic!("array is already locked"),
69 };
70 let array = std::mem::take(&mut *lock);
71 let array = array.into_shape_clone(shape).expect("invalid shape");
72 let _ = std::mem::replace(&mut *lock, array);
73 }
74
75 fn swap_axes(&mut self, axis_1: usize, axis_2: usize) {
76 let mut lock = match self.try_write() {
77 Ok(lock) => lock,
78 Err(TryLockError::Poisoned(_)) => panic!("array lock is poisoned"),
79 Err(TryLockError::WouldBlock) => panic!("array is already locked"),
80 };
81 lock.swap_axes(axis_1, axis_2);
82 }
83
84 fn move_data(
85 &mut self,
86 input: &dyn Array,
87 movements: &[mts_data_movement_t],
88 ) {
89 use ndarray::{Axis, Slice};
90
91 let input = input.as_any().downcast_ref::<Self>().expect("input must be a ndarray of the same type");
92 let input = match input.try_read() {
93 Ok(lock) => lock,
94 Err(TryLockError::Poisoned(_)) => panic!("input array lock is poisoned"),
95 Err(TryLockError::WouldBlock) => panic!("input array is already locked"),
96 };
97
98 let mut output = match self.try_write() {
99 Ok(lock) => lock,
100 Err(TryLockError::Poisoned(_)) => panic!("output array lock is poisoned"),
101 Err(TryLockError::WouldBlock) => panic!("output array is already locked"),
102 };
103
104 if movements.is_empty() {
105 return;
106 }
107
108 let first_prop_start_in = movements[0].properties_start_in;
110 let first_prop_start_out = movements[0].properties_start_out;
111 let first_prop_len = movements[0].properties_length;
112
113 let mut constant_properties = true;
114 let mut contiguous_input_samples = true;
115 let mut contiguous_output_samples = true;
116
117 for w in movements.windows(2) {
118 if w[0].properties_start_in != first_prop_start_in ||
119 w[0].properties_start_out != first_prop_start_out ||
120 w[0].properties_length != first_prop_len {
121 constant_properties = false;
122 break;
123 }
124
125 if w[1].sample_in != w[0].sample_in + 1 {
126 contiguous_input_samples = false;
127 }
128
129 if w[1].sample_out != w[0].sample_out + 1 {
130 contiguous_output_samples = false;
131 }
132 }
133
134 if constant_properties {
135 let last = movements.last().unwrap();
136 if last.properties_start_in != first_prop_start_in ||
137 last.properties_start_out != first_prop_start_out ||
138 last.properties_length != first_prop_len {
139 constant_properties = false;
140 }
141 }
142
143 let property_axis = output.shape().len() - 1;
144
145 if constant_properties {
146 let input_slice_info = Slice::from(first_prop_start_in..(first_prop_start_in + first_prop_len));
147 let output_slice_info = Slice::from(first_prop_start_out..(first_prop_start_out + first_prop_len));
148
149 if contiguous_input_samples && contiguous_output_samples {
150 let sample_start_in = movements[0].sample_in;
151 let sample_start_out = movements[0].sample_out;
152 let sample_count = movements.len();
153
154 let input_samples = input.slice_axis(
155 Axis(0),
156 Slice::from(sample_start_in..(sample_start_in + sample_count))
157 );
158 let mut output_samples = output.slice_axis_mut(
159 Axis(0),
160 Slice::from(sample_start_out..(sample_start_out + sample_count))
161 );
162
163 let value = input_samples.slice_axis(Axis(property_axis), input_slice_info);
164 let mut output_location = output_samples.slice_axis_mut(Axis(property_axis), output_slice_info);
165
166 output_location.assign(&value);
167 } else {
168 for move_item in movements {
169 let input_sample = input.index_axis(Axis(0), move_item.sample_in);
170 let mut output_sample = output.index_axis_mut(Axis(0), move_item.sample_out);
171
172 let value = input_sample.slice_axis(
173 Axis(property_axis - 1),
176 input_slice_info
177 );
178 let mut output_location = output_sample.slice_axis_mut(
179 Axis(property_axis - 1),
180 output_slice_info
181 );
182 output_location.assign(&value);
183 }
184 }
185 } else {
186 for move_item in movements {
188 let input_sample = input.index_axis(Axis(0), move_item.sample_in);
189 let mut output_sample = output.index_axis_mut(Axis(0), move_item.sample_out);
190
191 let value = input_sample.slice_axis(
192 Axis(property_axis - 1),
194 Slice::from(move_item.properties_start_in..(move_item.properties_start_in + move_item.properties_length))
195 );
196 let mut output_location = output_sample.slice_axis_mut(
197 Axis(property_axis - 1),
198 Slice::from(move_item.properties_start_out..(move_item.properties_start_out + move_item.properties_length))
199 );
200 output_location.assign(&value);
201 }
202 }
203 }
204
205 fn device(&self) -> DLDevice {
206 DLDevice::cpu()
207 }
208
209 fn dtype(&self) -> DLDataType {
210 T::get_dlpack_data_type()
211 }
212
213 fn as_dlpack(
214 &self,
215 device: DLDevice,
216 stream: Option<i64>,
217 max_version: DLPackVersion,
218 ) -> Result<DLPackTensor, Error> {
219 if stream.is_some() {
220 return Err(Error {
222 code: Some(crate::c_api::MTS_INVALID_PARAMETER_ERROR),
223 message: "CPU arrays can not be used with a stream".into(),
224 });
225 }
226 let vendored_version = DLPackVersion::current();
227 if max_version.major != vendored_version.major {
228 return Err(Error {
229 code: Some(crate::c_api::MTS_INVALID_PARAMETER_ERROR),
230 message: format!(
231 "invalid `max_version` in ndarray::ArrayD<T>::as_dlpack: \
232 we got v{}.{}, but we support v{}.{}",
233 max_version.major, max_version.minor,
234 vendored_version.major, vendored_version.minor
235 ),
236 });
237 }
238
239 let ndarray_device = DLDevice::cpu();
240
241 if device.device_type != ndarray_device.device_type || device.device_id != ndarray_device.device_id {
242 return Err(Error {
243 code: Some(crate::c_api::MTS_INVALID_PARAMETER_ERROR),
244 message: format!(
245 "Requested DLPack device ({}) does not match array device ({})",
246 device, ndarray_device
247 ),
248 });
249 }
250
251 let tensor: DLPackTensor = ReadOnly(Arc::clone(self)).try_into().map_err(|e| Error {
252 code: Some(crate::c_api::MTS_INVALID_PARAMETER_ERROR),
253 message: format!("failed to convert ndarray to DLPack: {:?}", e),
254 })?;
255
256 Ok(tensor)
257 }
258
259 #[allow(clippy::enum_glob_use)]
260 fn from_dlpack(&self, dlpack_tensor: DLPackTensor) -> Result<Box<dyn Array>, Error> {
261 use DLDataTypeCode::*;
262
263 let dtype = dlpack_tensor.dtype();
264
265 if dtype.lanes != 1 {
266 return Err(Error {
267 code: Some(crate::c_api::MTS_INVALID_PARAMETER_ERROR),
268 message: "Only DLPack tensors with lanes == 1 are supported".into(),
269 });
270 }
271
272 let map_error = |e| Error {
273 code: Some(crate::c_api::MTS_INVALID_PARAMETER_ERROR),
274 message: format!("failed to convert DLPack to ndarray: {:?}", e),
275 };
276
277 if dtype.code == kDLFloat && dtype.bits == 64 {
278 let array: ndarray::ArrayD<f64> = dlpack_tensor.try_into().map_err(map_error)?;
279 return Ok(Box::new(Arc::new(RwLock::new(array))));
280 } else if dtype.code == kDLFloat && dtype.bits == 32 {
281 let array: ndarray::ArrayD<f32> = dlpack_tensor.try_into().map_err(map_error)?;
282 return Ok(Box::new(Arc::new(RwLock::new(array))));
283 } else if dtype.code == kDLInt && dtype.bits == 8 {
284 let array: ndarray::ArrayD<i8> = dlpack_tensor.try_into().map_err(map_error)?;
285 return Ok(Box::new(Arc::new(RwLock::new(array))));
286 } else if dtype.code == kDLInt && dtype.bits == 16 {
287 let array: ndarray::ArrayD<i16> = dlpack_tensor.try_into().map_err(map_error)?;
288 return Ok(Box::new(Arc::new(RwLock::new(array))));
289 } else if dtype.code == kDLInt && dtype.bits == 32 {
290 let array: ndarray::ArrayD<i32> = dlpack_tensor.try_into().map_err(map_error)?;
291 return Ok(Box::new(Arc::new(RwLock::new(array))));
292 } else if dtype.code == kDLInt && dtype.bits == 64 {
293 let array: ndarray::ArrayD<i64> = dlpack_tensor.try_into().map_err(map_error)?;
294 return Ok(Box::new(Arc::new(RwLock::new(array))));
295 } else if dtype.code == kDLUInt && dtype.bits == 8 {
296 let array: ndarray::ArrayD<u8> = dlpack_tensor.try_into().map_err(map_error)?;
297 return Ok(Box::new(Arc::new(RwLock::new(array))));
298 } else if dtype.code == kDLUInt && dtype.bits == 16 {
299 let array: ndarray::ArrayD<u16> = dlpack_tensor.try_into().map_err(map_error)?;
300 return Ok(Box::new(Arc::new(RwLock::new(array))));
301 } else if dtype.code == kDLUInt && dtype.bits == 32 {
302 let array: ndarray::ArrayD<u32> = dlpack_tensor.try_into().map_err(map_error)?;
303 return Ok(Box::new(Arc::new(RwLock::new(array))));
304 } else if dtype.code == kDLUInt && dtype.bits == 64 {
305 let array: ndarray::ArrayD<u64> = dlpack_tensor.try_into().map_err(map_error)?;
306 return Ok(Box::new(Arc::new(RwLock::new(array))));
307 } else if dtype.code == kDLBool && dtype.bits == 8 {
308 let array: ndarray::ArrayD<bool> = dlpack_tensor.try_into().map_err(map_error)?;
309 return Ok(Box::new(Arc::new(RwLock::new(array))));
310 } else {
311 return Err(Error {
312 code: Some(crate::c_api::MTS_INVALID_PARAMETER_ERROR),
313 message: format!("Unsupported DLPack dtype {}", dtype),
314 });
315 }
316 }
317}
318
319#[cfg(test)]
320mod tests {
321 use dlpk::{DLPackPointerCast, GetDLPackDataType, sys::{DLDataTypeCode, DLDevice, DLPackVersion}};
322 use crate::MtsArray;
323
324 #[test]
325 fn ndarray_as_mts_array() {
326 let data = ndarray::Array::<f64, _>::zeros(vec![2, 3, 4]);
327 let mts_array = MtsArray::from(data);
328
329 assert_eq!(mts_array.shape().unwrap(), [2, 3, 4]);
330
331 let fill_value = MtsArray::from(ndarray::Array::from_elem(vec![], 42.0));
332
333 let created = mts_array.create(&[2, 3, 4], fill_value.as_ref()).unwrap();
334 assert_eq!(created.shape().unwrap(), [2, 3, 4]);
335 }
336
337 #[test]
338 fn ndarray_as_mts_array_dlpack() {
339 let data = ndarray::Array::<f64, _>::zeros(vec![4, 5, 6]);
340 let mts_array = MtsArray::from(data);
341
342 let dl_managed = mts_array.as_dlpack(DLDevice::cpu(), None, DLPackVersion::current()).unwrap();
343
344 assert_eq!(dl_managed.n_dims(), 3);
345 assert_eq!(dl_managed.shape(), [4, 5, 6]);
346
347 assert_eq!(dl_managed.dtype().code, DLDataTypeCode::kDLFloat);
348 assert_eq!(dl_managed.dtype().bits, 64);
349 assert_eq!(dl_managed.dtype().lanes, 1);
350 }
351
352 #[test]
353 fn ndarray_all_dtypes() {
354 fn test_for_dtype<T>(code: DLDataTypeCode, bits: u8) where T: Send + Sync + Clone + Default + GetDLPackDataType + DLPackPointerCast + 'static {
355 let data = ndarray::Array::<T, _>::from_elem(vec![2, 2], T::default());
356 let mts_array = MtsArray::from(data);
357
358 assert_eq!(mts_array.shape().unwrap(), [2, 2]);
359
360 let dl_managed = mts_array.as_dlpack(DLDevice::cpu(), None, DLPackVersion::current()).unwrap();
362 assert_eq!(dl_managed.dtype().code, code);
363 assert_eq!(dl_managed.dtype().bits, bits);
364 assert_eq!(dl_managed.dtype().lanes, 1);
365
366
367 let fill_value = MtsArray::from(ndarray::Array::from_elem(vec![], T::default()));
369
370 let created = mts_array.create(&[1, 1], fill_value.as_ref()).unwrap();
371 let dl_managed = created.as_dlpack(DLDevice::cpu(), None, DLPackVersion::current()).unwrap();
372
373 assert_eq!(dl_managed.dtype().code, code);
374 assert_eq!(dl_managed.dtype().bits, bits);
375 assert_eq!(dl_managed.dtype().lanes, 1);
376 }
377
378 test_for_dtype::<bool>(DLDataTypeCode::kDLBool, 8);
379 test_for_dtype::<f64>(DLDataTypeCode::kDLFloat, 64);
380 test_for_dtype::<f32>(DLDataTypeCode::kDLFloat, 32);
381 test_for_dtype::<i8>(DLDataTypeCode::kDLInt, 8);
382 test_for_dtype::<i16>(DLDataTypeCode::kDLInt, 16);
383 test_for_dtype::<i32>(DLDataTypeCode::kDLInt, 32);
384 test_for_dtype::<i64>(DLDataTypeCode::kDLInt, 64);
385 test_for_dtype::<u8>(DLDataTypeCode::kDLUInt, 8);
386 test_for_dtype::<u16>(DLDataTypeCode::kDLUInt, 16);
387 test_for_dtype::<u32>(DLDataTypeCode::kDLUInt, 32);
388 test_for_dtype::<u64>(DLDataTypeCode::kDLUInt, 64);
389 }
390
391 #[test]
392 fn ndarray_device() {
393 let data = ndarray::Array::<f64, _>::zeros(vec![2, 3]);
394 let mts_array = MtsArray::from(data);
395
396 assert_eq!(mts_array.device().unwrap(), DLDevice::cpu());
397 }
398
399 #[test]
400 fn as_dlpack_rejects_stream() {
401 let data = ndarray::Array::<f64, _>::zeros(vec![2, 3]);
402 let mts_array = MtsArray::from(data);
403 match mts_array.as_dlpack(DLDevice::cpu(), Some(42), DLPackVersion::current()) {
404 Err(e) => assert!(e.message.contains("stream"), "{}", e.message),
405 Ok(_) => panic!("expected error for non-null stream"),
406 }
407 }
408
409 #[test]
410 fn as_dlpack_rejects_wrong_device() {
411 let data = ndarray::Array::<f64, _>::zeros(vec![2, 3]);
412 let mts_array = MtsArray::from(data);
413 let cuda = DLDevice {
414 device_type: dlpk::sys::DLDeviceType::kDLCUDA,
415 device_id: 0,
416 };
417 match mts_array.as_dlpack(cuda, None, DLPackVersion::current()) {
418 Err(e) => assert!(e.message.contains("does not match"), "{}", e.message),
419 Ok(_) => panic!("expected error for CUDA device on CPU array"),
420 }
421 }
422
423 #[test]
424 fn as_dlpack_rejects_incompatible_version() {
425 let data = ndarray::Array::<f64, _>::zeros(vec![2, 3]);
426 let mts_array = MtsArray::from(data);
427
428 let bad_version = DLPackVersion { major: 99, minor: 0 };
429 match mts_array.as_dlpack(DLDevice::cpu(), None, bad_version) {
430 Err(e) => assert!(e.message.contains("version"), "{}", e.message),
431 Ok(_) => panic!("expected error for incompatible DLPack version"),
432 }
433 }
434
435 #[test]
436 #[allow(clippy::float_cmp)]
437 fn from_dlpack() {
438 let mut f64_data = ndarray::Array::<f64, _>::zeros(vec![2, 3]);
439 f64_data[[0, 0]] = 1.573;
440 f64_data[[1, 2]] = -42.0;
441 let f64_array = MtsArray::from(f64_data);
442
443 let mut i16_data = ndarray::Array::<i16, _>::zeros(vec![2, 5, 10]);
444 i16_data[[0, 1, 3]] = 3;
445 i16_data[[1, 2, 4]] = -42;
446 let i16_array = MtsArray::from(i16_data);
447
448 let f64_dl_tensor = f64_array.as_dlpack(DLDevice::cpu(), None, DLPackVersion::current()).unwrap();
449 let i16_dl_tensor = i16_array.as_dlpack(DLDevice::cpu(), None, DLPackVersion::current()).unwrap();
450
451 let new_f64_array = f64_array.from_dlpack(f64_dl_tensor).unwrap();
452 let new_i16_array = i16_array.from_dlpack(i16_dl_tensor).unwrap();
453
454 assert_eq!(f64_array.origin().unwrap(), i16_array.origin().unwrap());
455 assert_eq!(new_f64_array.origin().unwrap(), f64_array.origin().unwrap());
456 assert_eq!(new_i16_array.origin().unwrap(), i16_array.origin().unwrap());
457
458 let new_f64_dl_tensor = new_f64_array.as_dlpack(DLDevice::cpu(), None, DLPackVersion::current()).unwrap();
459 let new_i16_dl_tensor = new_i16_array.as_dlpack(DLDevice::cpu(), None, DLPackVersion::current()).unwrap();
460
461 let new_f64_data: ndarray::ArrayD<f64> = new_f64_dl_tensor.try_into().unwrap();
462 let new_i16_data: ndarray::ArrayD<i16> = new_i16_dl_tensor.try_into().unwrap();
463
464 assert_eq!(new_f64_data[[0, 0]], 1.573);
465 assert_eq!(new_f64_data[[1, 2]], -42.0);
466
467 assert_eq!(new_i16_data[[0, 1, 3]], 3);
468 assert_eq!(new_i16_data[[1, 2, 4]], -42);
469 }
470}