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

Unity中Texture2D转1x32x32x3 Tensor后ONNX模型预测异常求助

CIFAR-10模型Unity预测异常:所有输入均被判定为猫类(置信度约40%)

我在Unity中编写了一个脚本,接收Texture2D作为输入,用基于CIFAR-10训练的ONNX模型输出分类结果。之前针对MNIST的类似实现完全正常,所以问题应该出在Texture2D转Tensor的环节——这是两个脚本的唯一差异。

当前问题:无论传入什么图片,预测分数都极为接近,模型总是以约40%的置信度判定为猫类。我试过更换不同架构的训练模型、替换输入图片,也尝试直接用Texture构造Tensor,但都没解决问题,现在可以确定问题大概率在Texture2D转Tensor的过程中。

附上相关代码:

using System.Collections;
using System.Collections.Generic;
using System.Linq;
using UnityEngine;
using Unity.Barracuda;
using UI = UnityEngine.UI;
using TMPro;

public class Cifar10Script : MonoBehaviour
{
    public NNModel onnxAsset;
    public Texture2D imageToRecognise;
    public Texture2D tempImage;

    public UI.RawImage _imageView;
    public UI.Text _textView;

    public TMP_Text m_predictionTextDrawing;
    public TMP_Text m_predictionTextCamera;

    private int m_classPrediction;

    public void RunModel()
    {
        using Tensor input = new Tensor(1, 32,  32, 3);

        // Convert the input Texture2D into a 1x32x32x3 tensor.
        for (int y = 0; y < 32; y++)
        {
            for (int x = 0; x < 32; x++)
            {
                int tx = x * tempImage.width / 32;
                int ty = y * tempImage.height / 32;
                input[0, 31 - y, x, 0] = tempImage.GetPixel(tx, ty).r;
                input[0, 31 - y, x, 1] = tempImage.GetPixel(tx, ty).g;
                input[0, 31 - y, x, 2] = tempImage.GetPixel(tx, ty).b;
            }
        }

        IWorker worker = ModelLoader.Load(onnxAsset).CreateWorker(WorkerFactory.Device.CPU);

        worker.Execute(input);

        Tensor output = worker.PeekOutput();

        float[] scores = Enumerable.Range(0, 10).Select(i => output[i]).ToArray();

        float[] outputBuffer = output.ToReadOnlyArray();

        float lowValue = -999;
        int index = -1;

        for(int i = 0; i < 10; i++)
        {
            print(i + " " + outputBuffer[i]);
            if (outputBuffer[i] > lowValue)
            {
                lowValue = outputBuffer[i];
                index = i;
            }
        }

        worker.Dispose();

        if (outputBuffer[index] < 0.6f)
        {
            //m_predictionTextDrawing.SetText("?");
            //m_predictionTextCamera.SetText("?");
            m_classPrediction = -1;
            //return -1;
        }
        else
        {
            m_classPrediction = index;
            //m_predictionTextDrawing.SetText(index.ToString());
            //m_predictionTextCamera.SetText(index.ToString());
        }
        m_classPrediction = index;
    }

    public int GetClassPrediction()
    {
        return m_classPrediction;
    }
}

问题排查与修复建议

  • 图像缩放逻辑错误
    当前用整数除法做缩放的方式会导致图像严重失真,尤其是输入尺寸非32整数倍时。建议先将图片缩放到标准32x32尺寸再提取像素:

    // 先将tempImage缩放到32x32
    Texture2D scaledTexture = new Texture2D(32, 32);
    scaledTexture.SetPixels(tempImage.GetPixels());
    scaledTexture.Apply();
    // 后续直接使用scaledTexture的像素填充Tensor
    
  • 缺失像素归一化步骤
    CIFAR-10训练时通常会对输入做归一化处理(比如映射到[-1,1]区间),而UnityGetPixel返回的是[0,1]范围的颜色值,需和训练时的预处理对齐:

    Color pixel = scaledTexture.GetPixel(x, 31 - y);
    // 归一化到[-1,1](CIFAR-10常用预处理)
    input[0, y, x, 0] = pixel.r * 2 - 1;
    input[0, y, x, 1] = pixel.g * 2 - 1;
    input[0, y, x, 2] = pixel.b * 2 - 1;
    
  • Tensor维度顺序可能不匹配
    需确认训练模型时的输入维度顺序是[batch, height, width, channel]还是[batch, channel, height, width],如果是后者,需要调整通道索引位置:

    // 若模型期望CHW格式
    input[0, 0, y, x] = pixel.r * 2 - 1;
    input[0, 1, y, x] = pixel.g * 2 - 1;
    input[0, 2, y, x] = pixel.b * 2 - 1;
    
  • 模型输出未做概率转换
    很多ONNX模型最后一层是LogSoftmax或线性层,直接读取的输出不是置信度,需先做Softmax转换:

    // 将输出转换为概率值
    float sum = outputBuffer.Sum(Mathf.Exp);
    float[] probabilities = outputBuffer.Select(v => Mathf.Exp(v)/sum).ToArray();
    // 用probabilities数组查找最大置信度索引
    
  • 优化资源释放
    手动创建的scaledTexture需要在使用后销毁,避免内存泄漏:

    Destroy(scaledTexture);
    

内容的提问来源于stack exchange,提问作者Jbasham

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.24 16:14:57