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

sklearn模型转CoreML后,如何将UIImage转换为MLMultiArray?

把UIImage转换为MLMultiArray的解决方案

我完全懂你的困扰——用coremltools转换sklearn模型的时候,确实不像Keras那样能直接指定图像输入,只能用MLMultiArray作为输入,这就需要我们手动把UIImage转换成模型要求的格式。下面我给你一步步讲怎么实现,还有具体的Swift代码示例:

第一步:明确模型的输入要求

首先你得搞清楚CoreML模型期望的MLMultiArray的维度、数据类型和形状。比如你训练sklearn模型时用的是28x28的灰度图,那模型输入可能是[1, 28, 28](通道数、高度、宽度,CHW顺序),数据类型是float32。你可以用Xcode的模型检查器打开CoreML文件,确认这些参数,这一步很关键,不匹配的话模型会直接报错。

第二步:编写UIImage转MLMultiArray的Swift代码

我给你写一个UIImage的扩展方法,包含所有必要的预处理步骤(调整尺寸、颜色转换、归一化、数组转换):

import UIKit
import CoreML

extension UIImage {
    func convertToMLMultiArray(inputShape: [NSNumber], dataType: MLMultiArrayDataType = .float32) throws -> MLMultiArray {
        // 1. 调整图像尺寸到模型要求的大小
        guard let targetWidth = inputShape[2].intValue, let targetHeight = inputShape[1].intValue else {
            throw NSError(domain: "ImageConversionError", code: 0, userInfo: [NSLocalizedDescriptionKey: "Invalid input shape"])
        }
        let targetSize = CGSize(width: targetWidth, height: targetHeight)
        let resizedImage = resize(to: targetSize)
        
        // 2. 转换为灰度图(如果模型要求彩色输入,跳过这步转成RGB)
        guard let grayscaleImage = resizedImage.convertToGrayscale() else {
            throw NSError(domain: "ImageConversionError", code: 1, userInfo: [NSLocalizedDescriptionKey: "Failed to convert to grayscale"])
        }
        
        // 3. 获取图像像素数据并转换为Float数组(同步训练时的归一化逻辑)
        guard let pixelBuffer = grayscaleImage.toPixelBuffer() else {
            throw NSError(domain: "ImageConversionError", code: 2, userInfo: [NSLocalizedDescriptionKey: "Failed to get pixel buffer"])
        }
        
        var pixelValues = [Float]()
        CVPixelBufferLockBaseAddress(pixelBuffer, .readOnly)
        defer { CVPixelBufferUnlockBaseAddress(pixelBuffer, .readOnly) }
        
        guard let baseAddress = CVPixelBufferGetBaseAddress(pixelBuffer) else {
            throw NSError(domain: "ImageConversionError", code: 3, userInfo: [NSLocalizedDescriptionKey: "Failed to get pixel buffer base address"])
        }
        
        let bytesPerRow = CVPixelBufferGetBytesPerRow(pixelBuffer)
        for y in 0..<targetHeight {
            let row = baseAddress.advanced(by: y * bytesPerRow)
            for x in 0..<targetWidth {
                let pixel = row.advanced(by: x).load(as: UInt8.self)
                // 这里的归一化要和训练sklearn模型时的预处理完全一致,比如除以255
                pixelValues.append(Float(pixel) / 255.0)
            }
        }
        
        // 4. 创建MLMultiArray并填充数据
        let multiArray = try MLMultiArray(shape: inputShape, dataType: dataType)
        let pointer = multiArray.dataPointer.bindMemory(to: Float.self, capacity: multiArray.count)
        
        for i in 0..<multiArray.count {
            pointer[i] = pixelValues[i]
        }
        
        return multiArray
    }
    
    // 辅助方法:调整图像尺寸
    private func resize(to size: CGSize) -> UIImage {
        UIGraphicsBeginImageContextWithOptions(size, false, UIScreen.main.scale)
        defer { UIGraphicsEndImageContext() }
        draw(in: CGRect(origin: .zero, size: size))
        return UIGraphicsGetImageFromCurrentImageContext() ?? self
    }
    
    // 辅助方法:转换为灰度图
    private func convertToGrayscale() -> UIImage? {
        let context = CIContext(options: nil)
        guard let ciImage = CIImage(image: self) else { return nil }
        let grayscaleFilter = CIFilter(name: "CIColorControls")!
        grayscaleFilter.setValue(ciImage, forKey: kCIInputImageKey)
        grayscaleFilter.setValue(0.0, forKey: kCIInputSaturationKey)
        guard let outputImage = grayscaleFilter.outputImage else { return nil }
        let cgImage = context.createCGImage(outputImage, from: outputImage.extent)
        return UIImage(cgImage: cgImage!)
    }
    
    // 辅助方法:转换为CVPixelBuffer
    private func toPixelBuffer() -> CVPixelBuffer? {
        let attrs = [kCVPixelBufferCGImageCompatibilityKey: kCFBooleanTrue,
                     kCVPixelBufferCGBitmapContextCompatibilityKey: kCFBooleanTrue] as CFDictionary
        var pixelBuffer: CVPixelBuffer?
        let status = CVPixelBufferCreate(kCFAllocatorDefault,
                                         Int(size.width),
                                         Int(size.height),
                                         kCVPixelFormatType_OneComponent8,
                                         attrs,
                                         &pixelBuffer)
        guard status == kCVReturnSuccess, let buffer = pixelBuffer else { return nil }
        
        CVPixelBufferLockBaseAddress(buffer, [])
        defer { CVPixelBufferUnlockBaseAddress(buffer, []) }
        
        let context = CGContext(data: CVPixelBufferGetBaseAddress(buffer),
                                width: Int(size.width),
                                height: Int(size.height),
                                bitsPerComponent: 8,
                                bytesPerRow: CVPixelBufferGetBytesPerRow(buffer),
                                space: CGColorSpaceCreateDeviceGray(),
                                bitmapInfo: CGImageAlphaInfo.none.rawValue)
        context?.draw(cgImage!, in: CGRect(origin: .zero, size: size))
        return buffer
    }
}

第三步:使用转换后的MLMultiArray进行预测

你可以这样调用这个方法,然后传入模型进行预测:

// 假设你的CoreML模型叫MySklearnModel
guard let model = try? MySklearnModel(configuration: .init()) else {
    fatalError("Failed to load model")
}

// 假设模型输入形状是[1, 28, 28](通道数、高度、宽度)
guard let inputImage = UIImage(named: "test_image") else {
    fatalError("Failed to load image")
}

do {
    let inputArray = try inputImage.convertToMLMultiArray(inputShape: [1, 28, 28])
    let prediction = try model.prediction(input: inputArray)
    print("Prediction result: \(prediction)")
} catch {
    print("Error during conversion or prediction: \(error)")
}

关键注意事项

  • 维度顺序:一定要和CoreML模型的输入形状完全一致。比如如果你的sklearn模型训练时用的是HWC(高度、宽度、通道)顺序,那你需要调整转换后的数组顺序,改成模型要求的CHW或者其他格式。
  • 预处理一致性:训练模型时对图像做的所有预处理(比如归一化、颜色转换、裁剪),在转换UIImage时必须完全复刻,否则预测结果会不准确。
  • 数据类型:确保MLMultiArray的数据类型和模型要求的一致,比如模型要求double类型,就把代码里的Float改成Double。

内容的提问来源于stack exchange,提问作者Tom Shen

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.22 08:31:19