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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.18 07:19:57