如何保存Python机器学习模型以在C#编写的Unity游戏中调用
可行的模型序列化与跨平台加载方案
1. 使用ONNX格式
ONNX是跨框架的神经网络交换标准,几乎所有主流机器学习框架(scikit-learn、TensorFlow、PyTorch等)都支持导出为ONNX格式,Unity也有官方工具(Barracuda或ONNX Runtime)支持加载这类模型。
操作步骤:
Python端导出:以scikit-learn为例,借助
skl2onnx库转换模型:from skl2onnx import convert_sklearn from skl2onnx.common.data_types import FloatTensorType import onnx # 假设训练好的模型为clf,输入维度按需调整 initial_type = [('float_input', FloatTensorType([None, 4]))] onnx_model = convert_sklearn(clf, initial_types=initial_type) onnx.save_model(onnx_model, "model.onnx")PyTorch/TensorFlow可直接用框架自带的ONNX导出API完成转换。
Unity端加载推理:使用Unity Barracuda(轻量适配游戏场景):
using UnityEngine; using Unity.Barracuda; public class ModelRunner : MonoBehaviour { public NNModel onnxModel; private IWorker worker; void Start() { var model = ModelLoader.Load(onnxModel); worker = WorkerFactory.CreateWorker(WorkerFactory.Type.ComputePrecompiled, model); } public float[] Predict(float[] input) { var tensor = new Tensor(1, input.Length, input); worker.Execute(tensor); var output = worker.PeekOutput().ToArray(); tensor.Dispose(); return output; } }
2. 手动导出模型参数,C#重实现推理逻辑
如果模型结构简单(如线性回归、决策树),可直接导出模型核心参数(权重、系数、阈值等)为JSON/CSV等通用格式,再在C#中手动编写对应的推理代码。
操作步骤:
Python端导出参数:以线性回归为例:
import json model_params = { "coefficients": model.coef_.tolist(), "intercept": model.intercept_.tolist() } with open("model_params.json", "w") as f: json.dump(model_params, f)Unity端加载并推理:
using UnityEngine; using System.IO; using Newtonsoft.Json; public class LinearRegressionModel { public float[] coefficients; public float intercept; public static LinearRegressionModel Load(string path) { string json = File.ReadAllText(path); return JsonConvert.DeserializeObject<LinearRegressionModel>(json); } public float Predict(float[] input) { float result = intercept; for (int i = 0; i < coefficients.Length; i++) { result += coefficients[i] * input[i]; } return result; } }
3. 使用TensorFlow Lite格式
针对深度学习模型,TensorFlow Lite(TFLite)轻量化特性适配游戏场景,Unity有官方插件支持加载该格式模型。
操作步骤:
Python端导出TFLite模型:
import tensorflow as tf # 假设训练好的Keras模型为tf_model converter = tf.lite.TFLiteConverter.from_keras_model(tf_model) tflite_model = converter.convert() with open("model.tflite", "wb") as f: f.write(tflite_model)Unity端加载推理:
using UnityEngine; using TensorFlowLite; public class TFLiteRunner : MonoBehaviour { public TextAsset tfliteModel; private Interpreter interpreter; private float[] inputBuffer; private float[] outputBuffer; void Start() { interpreter = new Interpreter(tfliteModel.bytes); interpreter.AllocateTensors(); // 根据模型输入输出维度初始化缓冲 var inputShape = interpreter.GetInputTensorInfo(0).shape; var outputShape = interpreter.GetOutputTensorInfo(0).shape; inputBuffer = new float[inputShape[1]]; outputBuffer = new float[outputShape[1]]; } public float[] Predict(float[] input) { System.Array.Copy(input, inputBuffer, input.Length); interpreter.SetInputTensorData(0, inputBuffer); interpreter.Invoke(); interpreter.GetOutputTensorData(0, outputBuffer); return outputBuffer.Clone() as float[]; } }
内容的提问来源于stack exchange,提问作者RishtarCode47
相关产品推荐
相关产品推荐

