Swift中如何实现支持多数值类型的嵌套数组ShapedArray?
解决Swift中ShapedArray多维数值结构的类型转换问题
需求场景
想要在Swift中实现一个ShapedArray结构,用于处理N维数值数据,支持通过嵌套数组直接初始化,示例用法如下:
// 创建2x3的双精度二维数组 // ⎛ 1, 2, 3 ⎞ // ⎝ 4, 5, 6 ⎠ let arr = ShapedArray<Double>([[1, 2, 3], [4, 5, 6]]) // 创建2x2x3的单精度三维数组 // ⎛ ⎛ 1, 2, 3 ⎞ ⎞ // ⎜ ⎝ 4, 5, 6 ⎠ ⎟ // ⎜ ⎛ 7, 8, 9 ⎞ ⎟ // ⎝ ⎝ 0, 1, 2 ⎠ ⎠ let arr = ShapedArray<Float>([[[1, 2, 3], [4, 5, 6]], [[7, 8, 9], [0, 1, 2]]])
初始实现与问题
初始实现将数据存储为扁平化数组,通过getShape获取嵌套数组形状,flatten将嵌套数组转为一维数组:
func getShape(_ arr: some Collection) -> [Int] { if let first = arr.first as? any Collection { return [arr.count] + getShape(first) } else { return [arr.count] } } func flatten(_ arrays: [Any]) -> [Any] { var result = [Any]() for val in arrays { if let arr = val as? [Any] { result.append(contentsOf: flatten(arr)) } else { result.append(val) } } return result } struct ShapedArray<T> { let shape: [Int] let data: [T] init(arrays: [Any]) { self.shape = getShape(arrays) self.data = flatten(arrays) as! [T] } }
但该实现仅支持Int类型,创建Float或Double类型实例时会抛出类型转换错误:
Could not cast value of type 'Swift.Int' (0x7ff84ad1b2a0) to 'Swift.Float' (0x7ff84ad1ae28).
核心问题在于:
- 嵌套数组中的
Int字面量存入[Any]后,无法自动转换为指定的T类型 - 强制类型转换
as! [T]仅做直接类型匹配,不支持数值类型的自动转换
解决方案:类型安全的扁平化与转换
1. 重构形状计算逻辑
修改形状计算函数,支持泛型嵌套集合,避免模糊的类型转换:
// 处理嵌套集合 func getShape<C: Collection>(_ collection: C) -> [Int] where C.Element: Collection { var shape = [collection.count] if let firstElement = collection.first { shape += getShape(firstElement) } return shape } // 处理最底层非集合元素 func getShape<C: Collection>(_ collection: C) -> [Int] where C.Element: Any, C.Element: Collection == false { return [collection.count] }
2. 类型安全的扁平化转换函数
新增泛型函数,在遍历过程中尝试将元素转换为目标类型T,转换失败时抛出错误:
enum ShapedArrayError: Error { case invalidElementType case inconsistentShape } func flattenAndConvert<T>(_ value: Any, to type: T.Type) throws -> [T] where T: Numeric { if let element = value as? T { return [element] } else if let collection = value as? [Any] { var result = [T]() for item in collection { result.append(contentsOf: try flattenAndConvert(item, to: type)) } return result } else { throw ShapedArrayError.invalidElementType } }
3. 重构ShapedArray结构
修改初始化方法,增加形状一致性校验与错误处理,确保类型转换安全:
struct ShapedArray<T: Numeric> { let shape: [Int] let data: [T] init(_ nestedArray: Any) throws { guard isShapeConsistent(nestedArray) else { throw ShapedArrayError.inconsistentShape } self.shape = calculateShape(nestedArray) self.data = try flattenAndConvert(nestedArray, to: T.self) } // 校验嵌套数组形状是否一致 private func isShapeConsistent(_ value: Any) -> Bool { if let collection = value as? [Any] { guard !collection.isEmpty else { return true } let firstShape = calculateShape(collection[0]) return collection.allSatisfy { calculateShape($0) == firstShape } } else { return true } } // 统一计算形状 private func calculateShape(_ value: Any) -> [Int] { if let collection = value as? [Any] { guard !collection.isEmpty else { return [] } return [collection.count] + calculateShape(collection[0]) } else { return [] } } }
4. 使用示例
// 创建Double类型2D数组 do { let arr = try ShapedArray<Double>([[1, 2, 3], [4, 5, 6]]) print(arr.shape) // 输出: [2, 3] print(arr.data) // 输出: [1.0, 2.0, 3.0, 4.0, 5.0, 6.0] } catch { print(error) } // 创建Float类型3D数组 do { let arr = try ShapedArray<Float>([[[1, 2, 3], [4, 5, 6]], [[7, 8, 9], [0, 1, 2]]]) print(arr.shape) // 输出: [2, 2, 3] print(arr.data) // 输出: [1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 0.0, 1.0, 2.0] } catch { print(error) }
额外优化建议
- 添加多维下标访问器,支持
arr[0, 1]这类直观的多维索引 - 为
ShapedArray实现ExpressibleByArrayLiteral协议,简化初始化语法 - 增强错误提示,明确指出不规则嵌套数组的具体问题
内容的提问来源于stack exchange,提问作者wigging
相关产品推荐
相关产品推荐

