You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.29 07:40:41