You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

机器学习框架模型保存方式及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定义生成模型:

  1. 引入ONNX的Protobuf定义:将ONNX的onnx.proto文件转换为C#类(用protoc工具或Google.Protobuf NuGet包处理)
  2. 按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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.12 19:25:30