如何在C程序中加载Python(如SciKit-Learn)训练的机器学习模型?
在C中加载Python训练的机器学习模型的可行方案
当然有多种实用的C库和工具能实现这个需求,下面是几种主流方案:
1. TensorFlow Lite C API
适用于TensorFlow/Keras训练的模型,步骤如下:
- Python端:用
tf.lite.TFLiteConverter将训练好的SavedModel或Keras模型转换为.tflite格式。 - C端:调用TensorFlow Lite的C API加载模型、创建解释器并执行推理。
示例C代码片段:
#include "tensorflow/lite/c/c_api.h" int main() { // 加载.tflite模型文件 const TfLiteModel* model = TfLiteModelCreateFromFile("trained_model.tflite"); // 创建解释器配置 TfLiteInterpreterOptions* options = TfLiteInterpreterOptionsCreate(); // 初始化解释器 TfLiteInterpreter* interpreter = TfLiteInterpreterCreate(model, options); // 为张量分配内存 TfLiteInterpreterAllocateTensors(interpreter); // 获取输入输出张量指针 const TfLiteTensor* input_tensor = TfLiteInterpreterGetInputTensor(interpreter, 0); TfLiteTensor* output_tensor = TfLiteInterpreterGetOutputTensor(interpreter, 0); // 填充输入数据(示例为3维特征) float input_data[] = {0.5f, 1.2f, 3.1f}; TfLiteTensorCopyFromBuffer(input_tensor, input_data, sizeof(input_data)); // 执行推理 TfLiteInterpreterInvoke(interpreter); // 读取输出结果 float prediction_result[1]; TfLiteTensorCopyToBuffer(output_tensor, prediction_result, sizeof(prediction_result)); // 清理资源 TfLiteInterpreterDelete(interpreter); TfLiteInterpreterOptionsDelete(options); TfLiteModelDelete(model); return 0; }
2. ONNX Runtime C API
这是兼容性极强的方案,支持PyTorch、Scikit-learn、TensorFlow等多数框架训练的模型:
- Python端:将模型导出为ONNX格式(比如PyTorch用
torch.onnx.export,Scikit-learn用skl2onnx库转换)。 - C端:使用ONNX Runtime的C API创建会话,加载
.onnx模型并执行推理。
以Scikit-learn模型为例,Python导出代码:
from sklearn.ensemble import RandomForestClassifier from skl2onnx import convert_sklearn from skl2onnx.common.data_types import FloatTensorType # 假设已训练好随机森林模型 model = RandomForestClassifier() model.fit(X_train, y_train) # 转换为ONNX格式 initial_type = [("input_features", FloatTensorType([None, X_train.shape[1]]))] onnx_model = convert_sklearn(model, initial_types=initial_type) with open("rf_model.onnx", "wb") as f: f.write(onnx_model.SerializeToString())
C端通过ONNX Runtime加载该文件即可完成推理。
3. LightGBM/XGBoost原生C API
如果你的模型是用LightGBM或XGBoost训练的,这两个库本身提供原生C API,无需格式转换:
- Python端:直接保存模型(比如XGBoost用
model.save_model("xgb_model.model"))。 - C端:调用对应库的C API加载模型并预测。
示例XGBoost C代码片段:
#include <xgboost/c_api.h> #include <stdio.h> int main() { BoosterHandle booster; // 创建Booster并加载模型 XGBoosterCreate(NULL, 0, &booster); XGBoosterLoadModel(booster, "xgb_model.model"); // 准备输入数据(1个样本,3个特征) float input_data[] = {0.8f, 2.3f, 1.1f}; unsigned int shape[2] = {1, 3}; DMatrixHandle input_matrix; XGDMatrixCreateFromMat(input_data, shape[0], shape[1], 0.0f, &input_matrix); // 执行预测 bst_ulong pred_len; const float* preds; XGBoosterPredict(booster, input_matrix, 0, 0, 0, &pred_len, &preds); printf("预测结果:%f\n", preds[0]); // 清理资源 XGDMatrixFree(input_matrix); XGBoosterFree(booster); return 0; }
内容的提问来源于stack exchange,提问作者AndreGraveler
相关产品推荐
相关产品推荐

