Skip to main content

metatensor/data/
ndarray_array.rs

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        // Validate fill_value shape from the DLPack tensor directly
40        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        // Check if we can use the optimized path (all moves have same property structure)
109        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                        // property_axis - 1 because we are slicing the sample
174                        // axis out, so the property axis is now one less
175                        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            // fallback to the general case
187            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                    // see above for property_axis - 1 explanation
193                    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            // we only support CPU for now
221            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            // Should be able to export as DLPack
361            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            // And `create` should make an array of the same type (i32)
368            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}