@@ -423,6 +423,7 @@ where
423423 ctx. array . as_ptr ( )
424424 } ;
425425
426+ let ndim = ctx. shape . len ( ) as i32 ;
426427 let dl_tensor = sys:: DLTensor {
427428 // Casting to a mut pointer is not necessarily safe, but is required
428429 // by DLPack. The data can be mutated through this pointer, we
@@ -433,10 +434,10 @@ where
433434 device_type : sys:: DLDeviceType :: kDLCPU,
434435 device_id : 0 ,
435436 } ,
436- ndim : ctx . shape . len ( ) as i32 ,
437+ ndim : ndim ,
437438 dtype : T :: get_dlpack_data_type ( ) ,
438- shape : ctx. shape . as_mut_ptr ( ) ,
439- strides : ctx. strides . as_mut_ptr ( ) ,
439+ shape : if ndim == 0 { std :: ptr :: null_mut ( ) } else { ctx. shape . as_mut_ptr ( ) } ,
440+ strides : if ndim == 0 { std :: ptr :: null_mut ( ) } else { ctx. strides . as_mut_ptr ( ) } ,
440441 byte_offset : 0 ,
441442 } ;
442443
@@ -489,8 +490,8 @@ where
489490 } ,
490491 ndim,
491492 dtype : T :: get_dlpack_data_type ( ) ,
492- shape : ctx. shape . as_mut_ptr ( ) ,
493- strides : ctx. strides . as_mut_ptr ( ) ,
493+ shape : if ndim == 0 { std :: ptr :: null_mut ( ) } else { ctx. shape . as_mut_ptr ( ) } ,
494+ strides : if ndim == 0 { std :: ptr :: null_mut ( ) } else { ctx. strides . as_mut_ptr ( ) } ,
494495 byte_offset : 0 ,
495496 } ;
496497
@@ -786,7 +787,7 @@ mod tests {
786787 assert_eq ! ( array_view. shape( ) , & [ 0 , 0 , 0 ] ) ;
787788 }
788789
789- unsafe extern "C" fn empty_deleter ( tensor : * mut sys:: DLManagedTensorVersioned ) {
790+ unsafe extern "C" fn box_deleter ( tensor : * mut sys:: DLManagedTensorVersioned ) {
790791 let _ = Box :: from_raw ( tensor) ;
791792 }
792793
@@ -811,7 +812,7 @@ mod tests {
811812 let managed = Box :: new ( crate :: sys:: DLManagedTensorVersioned {
812813 version : crate :: sys:: DLPackVersion :: current ( ) ,
813814 manager_ctx : std:: ptr:: null_mut ( ) ,
814- deleter : Some ( empty_deleter ) ,
815+ deleter : Some ( box_deleter ) ,
815816 flags : 0 ,
816817 dl_tensor,
817818 } ) ;
@@ -820,4 +821,79 @@ mod tests {
820821 let array: Array3 < f32 > = tensor. try_into ( ) . unwrap ( ) ;
821822 assert_eq ! ( array. shape( ) , & [ 0 , 0 , 0 ] ) ;
822823 }
824+
825+ #[ test]
826+ fn scalar_ndarray_to_dlpack ( ) {
827+ let array = arr0 ( 42.0f64 ) ;
828+ let tensor: DLPackTensor = array. try_into ( ) . unwrap ( ) ;
829+ assert_eq ! ( tensor. n_dims( ) , 0 ) ;
830+ assert ! ( tensor. as_dltensor( ) . shape. is_null( ) ) ;
831+ assert ! ( tensor. shape( ) . is_empty( ) ) ;
832+ assert ! ( tensor. as_dltensor( ) . strides. is_null( ) ) ;
833+ assert ! ( tensor. strides( ) . is_none( ) ) ;
834+ }
835+
836+ #[ test]
837+ fn scalar_arc_array_to_dlpack ( ) {
838+ let array = ndarray:: ArcArray :: < f64 , ndarray:: Ix0 > :: from_elem ( ( ) , 42.0f64 ) ;
839+ let tensor: DLPackTensor = array. try_into ( ) . unwrap ( ) ;
840+ assert_eq ! ( tensor. n_dims( ) , 0 ) ;
841+ assert ! ( tensor. as_dltensor( ) . shape. is_null( ) ) ;
842+ assert ! ( tensor. shape( ) . is_empty( ) ) ;
843+ assert ! ( tensor. as_dltensor( ) . strides. is_null( ) ) ;
844+ assert ! ( tensor. strides( ) . is_none( ) ) ;
845+ }
846+
847+ #[ test]
848+ fn scalar_dlpack_to_ndarray_view ( ) {
849+ let mut value = 3.41f32 ;
850+
851+ let dl_tensor = DLTensor {
852+ data : ( & mut value as * mut f32 ) . cast ( ) ,
853+ device : DLDevice {
854+ device_type : DLDeviceType :: kDLCPU,
855+ device_id : 0 ,
856+ } ,
857+ ndim : 0 ,
858+ dtype : f32:: get_dlpack_data_type ( ) ,
859+ shape : std:: ptr:: null_mut ( ) ,
860+ strides : std:: ptr:: null_mut ( ) ,
861+ byte_offset : 0 ,
862+ } ;
863+
864+ let dlpack_ref = unsafe { DLPackTensorRef :: from_raw ( dl_tensor) } ;
865+ let array_view = ArrayView0 :: < f32 > :: try_from ( dlpack_ref) . unwrap ( ) ;
866+ assert ! ( array_view. shape( ) . is_empty( ) ) ;
867+ assert_eq ! ( array_view[ ( ) ] , 3.41 ) ;
868+ }
869+
870+ #[ test]
871+ fn scalar_dlpack_to_ndarray_owned ( ) {
872+ let mut value = 2.72f64 ;
873+
874+ let dl_tensor = DLTensor {
875+ data : ( & mut value as * mut f64 ) . cast ( ) ,
876+ device : DLDevice {
877+ device_type : DLDeviceType :: kDLCPU,
878+ device_id : 0 ,
879+ } ,
880+ ndim : 0 ,
881+ dtype : f64:: get_dlpack_data_type ( ) ,
882+ shape : std:: ptr:: null_mut ( ) ,
883+ strides : std:: ptr:: null_mut ( ) ,
884+ byte_offset : 0 ,
885+ } ;
886+
887+ let managed = Box :: new ( crate :: sys:: DLManagedTensorVersioned {
888+ version : crate :: sys:: DLPackVersion :: current ( ) ,
889+ manager_ctx : std:: ptr:: null_mut ( ) ,
890+ deleter : Some ( box_deleter) ,
891+ flags : 0 ,
892+ dl_tensor,
893+ } ) ;
894+
895+ let tensor = unsafe { DLPackTensor :: from_ptr ( Box :: into_raw ( managed) ) } ;
896+ let array: Array0 < f64 > = tensor. try_into ( ) . unwrap ( ) ;
897+ assert_eq ! ( array[ ( ) ] , 2.72 ) ;
898+ }
823899}
0 commit comments