如何将sklearn训练的Random Forest模型转换为.tflite格式
结论
你训练好的scikit-learn Random Forest模型可以转换为.tflite格式,需要先将sklearn格式的模型转换为中间格式,再导出为tflite,具体操作如下:
具体转换步骤
- 第一步:安装所需依赖包
pip install tensorflow skl2onnx onnxruntime onnx-tf joblib
- 第二步:在你现有训练代码的末尾,将训练好的随机森林模型转换为ONNX中间格式
import skl2onnx from skl2onnx import convert_sklearn from skl2onnx.common.data_types import FloatTensorType # 定义输入特征形状:batch维度不固定,特征数为你用到的6个气象特征 initial_type = [('float_input', FloatTensorType([None, 6]))] # 转换sklearn随机森林为ONNX格式 onnx_model = convert_sklearn(clf, initial_types=initial_type) # 保存ONNX模型文件 with open("rf_weather.onnx", "wb") as f: f.write(onnx_model.SerializeToString())
- 第三步:将ONNX模型转换为TensorFlow SavedModel格式
import onnx from onnx_tf.backend import prepare # 加载刚才导出的ONNX模型 onnx_model = onnx.load("rf_weather.onnx") # 转换为TensorFlow可识别的表示 tf_rep = prepare(onnx_model) # 导出为SavedModel文件夹 tf_rep.export_graph("tf_rf_model")
- 第四步:将SavedModel转换为最终的.tflite格式
import tensorflow as tf # 加载SavedModel创建转换器 converter = tf.lite.TFLiteConverter.from_saved_model("tf_rf_model") # 配置支持随机森林所需的算子 converter.target_spec.supported_ops = [ tf.lite.OpsSet.TFLITE_BUILTINS, tf.lite.OpsSet.SELECT_TF_OPS ] # 执行转换 tflite_model = converter.convert() # 保存tflite模型文件 with open("rf_weather.tflite", "wb") as f: f.write(tflite_model)
注意事项
- 你需要同时保存训练时用到的
StandardScaler和LabelEncoder组件,移动端推理时,输入的气象数据要先做和训练流程一致的标准化处理,推理输出结果还要用标签编码器做逆转换才能得到最终可识别的降雨相关结果,保存代码如下:
import joblib # 保存标准化器和标签编码器 joblib.dump(scaler, "scaler.joblib") joblib.dump(gender_encoder, "label_encoder.joblib")
- 如果转换后模型体积过大或者移动端推理速度慢,可以适当降低随机森林的树数量,比如把
n_estimators从100调整到30~70的区间,平衡精度和推理性能。 - 转换完成后建议先在本地对比tflite模型和原sklearn模型的推理结果,避免转换过程出现精度损失。
内容的提问来源于stack exchange,提问作者NishPre
相关产品推荐
相关产品推荐

