Scala中如何简洁获取多维数组的shape?类似NumPy的a.shape功能
在Scala中获取多维数组Shape的简洁方法
当然可以实现类似NumPy的功能啦!Scala里的多维数组本质是数组的嵌套(比如Array[Array[Int]]就是二维数组),我们可以根据这个特性来编写简洁的方法获取它的shape。下面分情况给你介绍实用方案:
1. 针对二维数组的快速实现
如果只是处理二维数组,直接取外层数组长度和第一个子数组的长度就可以,记得要处理空数组的边界情况:
// 定义一个二维数组 val a = Array(Array(1,2,3), Array(4,5,6)) // 获取shape,避免空数组报错 val shape = if (a.isEmpty) (0, 0) else (a.length, a.head.length) println(shape) // 输出 (2, 3)
2. 通用的任意维度数组Shape获取
如果需要支持任意维度的数组(比如三维、四维),可以用模式匹配+递归写一个通用函数,自动遍历每个维度的长度:
def getShape(arr: Any): List[Int] = arr match { // 非空数组:取当前维度长度,再递归获取子数组的维度 case array: Array[_] if array.nonEmpty => array.length :: getShape(array.head) // 空数组:只记录当前维度长度 case array: Array[_] => List(array.length) // 非数组元素(到达最内层):停止递归 case _ => Nil } // 测试二维数组 val a2d = Array(Array(1,2,3), Array(4,5,6)) println(getShape(a2d)) // 输出 List(2, 3) // 测试三维数组 val a3d = Array(Array(Array(1,2), Array(3,4)), Array(Array(5,6), Array(7,8))) println(getShape(a3d)) // 输出 List(2, 2, 2)
注意事项
和NumPy的规整数组不同,Scala的嵌套数组允许不规则维度(比如二维数组里有的子数组长度不一样)。上面的方法会以第一个子数组的维度为准,如果你的数组是不规则的,可能需要额外处理这种情况哦。
内容的提问来源于stack exchange,提问作者andrew
相关产品推荐
相关产品推荐

