如何在ML.NET中复用已训练的TensorFlow回归神经网络模型?
解决方案
1. 首先修正ONNX转换命令
你当前的转换命令存在两处核心错误:一是savedModels\和--output之间缺少空格,导致SavedModel路径读取异常,生成的ONNX模型本身损坏,这也是Netron无法识别模型的直接原因;二是没有固定输入输出命名,导致自动生成的节点名混乱无法匹配。
首先升级tf2onnx到最新版本:pip install -U tf2onnx
再执行正确的转换命令:python -m tf2onnx.convert --saved-model ./savedModels --output model_new.onnx --opset 15 --batch-size 1
添加--batch-size 1参数可以固定输入的batch维度为1,避免C#端维度不匹配的问题,--opset 15可以提升ONNX模型在ML.NET中的兼容性。
2. 手动确认ONNX输入输出名称
如果转换后的模型还是无法用Netron打开,直接通过Python代码读取模型结构即可拿到准确的输入输出信息:
import onnx model = onnx.load("model_new.onnx") # 打印所有输入信息 print("模型输入列表:") for input_node in model.graph.input: print(f"名称:{input_node.name}, 形状:{[dim.dim_value for dim in input_node.type.tensor_type.shape.dim]}") # 打印所有输出信息 print("模型输出列表:") for output_node in model.graph.output: print(f"名称:{output_node.name}, 形状:{[dim.dim_value for dim in output_node.type.tensor_type.shape.dim]}")
3. 修正C#端适配代码
将上一步查到的输入输出名称对应替换到C#代码中即可,参考修改后的代码如下:
class Program { static string ONNX_MODEL_PATH = @"C:\Users\FedorovEA\Downloads\model_new_9.onnx"; static void Main(string[] args) { MLContext mlContext = new MLContext(); var onnxPredictionPipeline = GetPredictionPipeline(mlContext); var testInput = new OnnxInput { // 按实际业务填入6个输入值 ModelInput = new float[] {5.0f, 3.0f, 200.0f, 200.0f, 10.0f, 5.0f} }; var onnxPredictionEngine = mlContext.Model.CreatePredictionEngine<OnnxInput, OnnxOutput>(onnxPredictionPipeline); var prediction = onnxPredictionEngine.Predict(testInput); Console.WriteLine($"Predicted Fare: {prediction.PredictedFare.First()}"); } static ITransformer GetPredictionPipeline(MLContext mlContext) { // 此处替换为你实际查询到的输入、输出节点名称 var inputColumns = new string[] {"serving_default_normalization_input"}; var outputColumns = new string[] { "dense_9" }; var onnxPredictionPipeline = mlContext.Transforms.ApplyOnnxModel(outputColumnNames: outputColumns, inputColumnNames: inputColumns, ONNX_MODEL_PATH); var emptyDv = mlContext.Data.LoadFromEnumerable(new OnnxInput[] { }); return onnxPredictionPipeline.Fit(emptyDv); } } public class OnnxInput { /// <summary> /// 6维浮点数输入 /// </summary> // 此处替换为实际查询到的输入节点名称 [ColumnName("serving_default_normalization_input")] public float[] ModelInput { get; set; } } public class OnnxOutput { // 此处替换为实际查询到的输出节点名称 [ColumnName("dense_9")] public float[] PredictedFare { get; set; } }
4. 兼容性备选方案
如果你使用的tf2onnx版本对Keras的Normalization层转换存在兼容问题,可以直接把归一化逻辑移到C#端实现,模型结构去掉normalizer层后再转换,可完全避免输入节点异常的问题。
内容的提问来源于stack exchange,提问作者Евгений Федоров
相关产品推荐
相关产品推荐

