将Sklearn随机森林回归模型转换为iOS可用的Core ML模型
Sklearn随机森林转Core ML 解决指南
一、input_features参数匹配说明
- 不需要和原特征矩阵X的特征名严格匹配,但必须保证特征顺序完全一致。Core ML模型推理时是按输入数组的顺序对应训练时的特征顺序,特征名仅用于标识,不影响计算逻辑。
- 当X是无特征标签的
np.ndarray时,你可以手动生成任意合法的特征名列表(比如["feature_0", "feature_1", ..., "feature_n"]),只要数量和X的特征维度一致即可。
二、无特征标签矩阵的模型转换步骤
假设你的Sklearn随机森林回归模型已训练完成,X为(样本数, 特征数)格式的np.ndarray,步骤如下:
- 生成输入特征名列表:根据X的特征维度批量生成名称
feature_count = X.shape[1] input_features = [f"feature_{i}" for i in range(feature_count)] - 转换并保存Core ML模型
import coremltools as ct from sklearn.ensemble import RandomForestRegressor # 假设rf_model是训练好的随机森林回归模型 rf_model = RandomForestRegressor() # 此处省略训练、缺失值插补步骤 # 执行转换 coreml_model = ct.converters.sklearn.convert( rf_model, input_features=input_features, output_feature_names=["esg_score"] ) # 保存模型文件 coreml_model.save("ESGPredictor.mlmodel") - 验证模型结构:打印
coreml_model可查看输入输出特征的维度、名称信息,确认参数匹配。
三、Swift端模型使用示例
1. 导入模型
将生成的ESGPredictor.mlmodel拖入Xcode项目,Xcode会自动生成同名Swift类。
2. 基础推理代码
import CoreML // 输入特征数组:顺序必须和训练时的特征顺序完全一致 let inputFeatures: [Double] = [0.2, 1.5, 3.0] // 替换为实际特征值 do { // 初始化模型 let model = try ESGPredictor(configuration: .init()) // 构造输入:按模型生成的特征名依次传入值 let input = ESGPredictorInput( feature_0: inputFeatures[0], feature_1: inputFeatures[1], feature_2: inputFeatures[2] // 特征数量多的话依次补全feature_n ) // 执行预测 let output = try model.prediction(input: input) // 获取ESG评分结果 let esgScore = output.esg_score print("预测ESG评分:\(esgScore)") } catch { print("模型加载/预测失败:\(error.localizedDescription)") }
3. 批量特征简化输入(可选)
如果特征数量较多,可通过反射简化输入构造:
import CoreML func predictESG(features: [Double]) -> Double? { // 先校验特征数量是否匹配 guard features.count == 你的特征总数 else { return nil } do { let model = try ESGPredictor(configuration: .init()) let inputNames = model.modelDescription.inputDescriptionsByName.keys var inputDict = [String: Any]() // 按顺序映射特征名和输入值 for (index, name) in inputNames.enumerated() { inputDict[name] = features[index] } let input = try ESGPredictorInput(dictionary: inputDict) let output = try model.prediction(input: input) return output.esg_score } catch { print("错误:\(error)") return nil } } // 使用示例 if let score = predictESG(features: [0.2, 1.5, 3.0]) { print("ESG评分:\(score)") }
内容的提问来源于stack exchange,提问作者Миршод Махсудов
相关产品推荐
相关产品推荐

