机器学习框架模型保存方式及C#自制CNN手动转ONNX方法问询
自研CNN转ONNX的手动实现指南
一、明确CNN需保存的核心数据结构
不管是Pickle还是ONNX,本质都是序列化模型的可学习参数与网络结构定义。你的自研CNN必须提取以下关键信息:
- 网络层的类型与执行顺序:比如Conv2D、ReLU、MaxPool、Dense的排列逻辑
- 各层的可训练参数:
- 卷积层:卷积核权重(
weight)、偏置(bias,若启用) - 全连接层:权重矩阵、偏置向量
- 池化/激活层:无训练参数,但需记录配置(如池化窗口大小、步幅)
- 卷积层:卷积核权重(
- 输入输出维度:比如输入图像的高/宽/通道数、输出类别数
C#参数结构示例
// 卷积层参数容器 public class ConvLayerParams { public int InChannels { get; set; } public int OutChannels { get; set; } public int KernelSize { get; set; } public int Stride { get; set; } public int Padding { get; set; } public NDarray Weights { get; set; } // Numpy.NET的NDarray类型 public NDarray? Bias { get; set; } } // 全连接层参数容器 public class DenseLayerParams { public int InFeatures { get; set; } public int OutFeatures { get; set; } public NDarray Weights { get; set; } public NDarray Bias { get; set; } } // 整个模型的结构与参数集合 public class CNNModel { public List<object> Layers { get; set; } // 存储各层参数实例(Conv/Dense等) public (int Channels, int Height, int Width) InputShape { get; set; } public int OutputClasses { get; set; } }
二、手动构建ONNX模型核心逻辑
ONNX基于Protobuf定义结构,C#中可直接借助ONNX的官方Protobuf定义生成模型:
- 引入ONNX的Protobuf定义:将ONNX的
onnx.proto文件转换为C#类(用protoc工具或Google.ProtobufNuGet包处理) - 按ONNX规范逐步构建Graph:
- 定义输入输出张量:指定名称、数据类型、维度
- 逐个添加网络层节点(Node):每个Node对应CNN的一层,比如Conv节点需指定输入输出名、属性(kernel_size、stride等),并将参数(权重、偏置)作为初始张量(Initializers)加入Graph
- 将所有可训练参数转换为ONNX的
TensorProto格式,加入Initializers列表
关键代码片段
// 初始化ONNX模型 var model = new ModelProto(); model.IrVersion = 8; // 对应ONNX IR版本,按需选择 model.ProducerName = "Custom-CSharp-CNN"; // 构建Graph var graph = new GraphProto(); model.Graph = graph; // 添加输入张量(示例:1x28x28单通道图像) var inputTensor = new ValueInfoProto(); inputTensor.Name = "input"; inputTensor.Type.TensorType.Shape.Dim.Add(new TensorShapeProto.Types.Dim { DimValue = 1 }); // Batch Size inputTensor.Type.TensorType.Shape.Dim.Add(new TensorShapeProto.Types.Dim { DimValue = 1 }); // Channels inputTensor.Type.TensorType.Shape.Dim.Add(new TensorShapeProto.Types.Dim { DimValue = 28 }); // Height inputTensor.Type.TensorType.Shape.Dim.Add(new TensorShapeProto.Types.Dim { DimValue = 28 }); // Width inputTensor.Type.TensorType.ElemType = (int)TensorProto.Types.DataType.Float; graph.Input.Add(inputTensor); // 添加输出张量(示例:10分类) var outputTensor = new ValueInfoProto(); outputTensor.Name = "output"; outputTensor.Type.TensorType.Shape.Dim.Add(new TensorShapeProto.Types.Dim { DimValue = 1 }); outputTensor.Type.TensorType.Shape.Dim.Add(new TensorShapeProto.Types.Dim { DimValue = 10 }); outputTensor.Type.TensorType.ElemType = (int)TensorProto.Types.DataType.Float; graph.Output.Add(outputTensor); // 遍历自研CNN的层,生成ONNX节点与初始张量 string currentInputName = "input"; int layerIndex = 0; foreach (var layer in cnnModel.Layers) { if (layer is ConvLayerParams convParams) { // 创建Conv节点 var convNode = new NodeProto(); convNode.Name = $"conv_{layerIndex}"; convNode.OpType = "Conv"; convNode.Input.Add(currentInputName); convNode.Input.Add($"conv_weights_{layerIndex}"); if (convParams.Bias != null) convNode.Input.Add($"conv_bias_{layerIndex}"); var outputName = $"conv_output_{layerIndex}"; convNode.Output.Add(outputName); // 设置Conv属性:kernel_size、stride、padding(ONNX的pads为[top, left, bottom, right]) convNode.Attributes["kernel_shape"].Ints.Add(convParams.KernelSize); convNode.Attributes["kernel_shape"].Ints.Add(convParams.KernelSize); convNode.Attributes["strides"].Ints.Add(convParams.Stride); convNode.Attributes["strides"].Ints.Add(convParams.Stride); convNode.Attributes["pads"].Ints.AddRange(new[] {convParams.Padding, convParams.Padding, convParams.Padding, convParams.Padding}); graph.Node.Add(convNode); // 转换卷积权重为TensorProto并加入Initializers var weightTensor = ConvertNdarrayToTensorProto(convParams.Weights, $"conv_weights_{layerIndex}"); graph.Initializer.Add(weightTensor); // 若有偏置,同样转换添加 if (convParams.Bias != null) { var biasTensor = ConvertNdarrayToTensorProto(convParams.Bias, $"conv_bias_{layerIndex}"); graph.Initializer.Add(biasTensor); } currentInputName = outputName; layerIndex++; } // ReLU/MaxPool/Dense层同理,对应ONNX的OpType分别为Relu、MaxPool、Gemm } // 保存ONNX文件 using (var stream = File.Create("custom_cnn.onnx")) { model.WriteTo(stream); } // 辅助函数:Numpy.NET NDarray转ONNX TensorProto private TensorProto ConvertNdarrayToTensorProto(NDarray ndarray, string name) { var tensor = new TensorProto(); tensor.Name = name; tensor.DataType = (int)TensorProto.Types.DataType.Float; // 设置张量维度 foreach (var dim in ndarray.Shape) { tensor.Dims.Add((long)dim); } // 将NDarray数据转为float数组写入TensorProto float[] data = ndarray.GetData<float>(); tensor.FloatData.AddRange(data); return tensor; }
三、验证ONNX模型正确性
保存后可通过ONNX Runtime(C#中用Microsoft.ML.OnnxRuntime NuGet包)加载模型,对比自研CNN的输出结果,确保参数与结构转换正确:
using var session = new InferenceSession("custom_cnn.onnx"); // 构造测试输入张量 var inputTensor = new DenseTensor<float>(new float[1*1*28*28], new[] {1,1,28,28}); // 填充测试数据... var inputs = new List<NamedOnnxValue> { NamedOnnxValue.CreateFromTensor("input", inputTensor) }; using var outputs = session.Run(inputs); var onnxOutput = outputs.First().AsTensor<float>(); // 与自研CNN的输出对比验证
替代简化方案
若手动构建ONNX过于繁琐,可先将模型参数序列化为二进制文件(比如用System.Text.Json,Numpy.NET的NDarray先转成float数组再序列化),后续再写工具将二进制参数转换为ONNX格式,但直接手动构建ONNX的兼容性更好。
内容的提问来源于stack exchange,提问作者Adotac
相关产品推荐
相关产品推荐

