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

ESP32CAM加载TFLite模型始终输出0.5分类分数求助

纸板/纸张与塑料分类模型部署ESP32CAM异常问题

问题描述

开发了一个基于MobileNetV2的分类模型,用于区分纸板/纸张与塑料。在Google Colab中加载TFLite模型运行正常,但部署到ESP32CAM(AI Thinker)后,模型对所有输入(包括无输入)的预测分数始终为0.5,输出结果完全一致,怀疑神经网络实现存在问题。

相关代码与日志

NeuralNetwork.cpp代码

#include "NeuralNetwork.h"
#include "model_160x120_data.h"
#include "tensorflow/lite/micro/all_ops_resolver.h"
#include "tensorflow/lite/micro/micro_error_reporter.h"
#include "tensorflow/lite/micro/micro_interpreter.h"
#include <Arduino.h>
#include "esp_camera.h"


float* model_input_buffer = nullptr;
const int kArenaSize = 1900000;
NeuralNetwork::NeuralNetwork()
{
    error_reporter = new tflite::MicroErrorReporter();
    model = tflite::GetModel(garbageTfLite160x120_tflite);

    
    // 加载所需算子实现
    resolver = new tflite::MicroMutableOpResolver<10>();

    resolver->AddAveragePool2D();
    resolver->AddConv2D();
    resolver->AddDepthwiseConv2D();
    resolver->AddReshape();
    resolver->AddSoftmax();
    resolver->AddAdd();
    resolver->AddPad();
    resolver->AddPadV2();
    resolver->AddMean();
    resolver->AddFullyConnected();


    tensor_arena = (uint8_t *) ps_malloc(kArenaSize); // ESP32CAM有PSRAM,使用ps_malloc(尝试过malloc,结果仍为0.5)
    if (!tensor_arena)
    {
        TF_LITE_REPORT_ERROR(error_reporter, "Could not allocate arena");
        return;
    }
    // 创建解释器运行模型
    interpreter = new tflite::MicroInterpreter(
        model, *resolver, tensor_arena, kArenaSize, error_reporter);
    // 为模型张量分配内存
    TfLiteStatus allocate_status = interpreter->AllocateTensors();
    if (allocate_status != kTfLiteOk)
    {
        TF_LITE_REPORT_ERROR(error_reporter, "AllocateTensors() failed");
        return;
    }
    size_t used_bytes = interpreter->arena_used_bytes();
    TF_LITE_REPORT_ERROR(error_reporter, "Used bytes by the model: %d\n", used_bytes);

    // 获取输入输出张量指针
    input = interpreter->input(0);
    model_input_buffer = input->data.f;
    output = interpreter->output(0);

    int batch_size = input->dims->data[0];
    int h = input->dims->data[1];
    int w = input->dims->data[2];
    int channels = input->dims->data[3];
    Serial.print("\n ");
    Serial.print(batch_size);
    Serial.print("\n ");
    Serial.print(h);
    Serial.print("\n ");
    Serial.print(w);
    Serial.print("\n ");
    Serial.print(channels);
    Serial.print("\n ");
}


float NeuralNetwork::classify_image(camera_fb_t *fb) {
 
    Serial.println ("DIMS: ");
    Serial.println(input->dims->size);
    Serial.println("input-> bytes: ");
    Serial.println(input->bytes);
    
    int img_size = 120*160*3;
    for (int i=0; i < img_size; i++) {
        model_input_buffer[i] = fb->buf[i]/255.0;
    }


    if (kTfLiteOk != interpreter->Invoke()) 
    {
        error_reporter->Report("Error");
    }
    
    TfLiteTensor* output = interpreter->output(0);
    Serial.println("output0: ");
    Serial.println(output->data.f[0]);
    Serial.println("output1: ");
    Serial.println(output->data.f[1]);
    Serial.println("output2: ");
    Serial.println(output->data.f[2]);
    return output->data.f[0];

}

设备输出日志

Used bytes by the model: 1898436

batch_size: 1

h: 120

w: 160

channels: 3

DIMS: 4
input-> bytes:230400

output0: 0.50

output1: 0.50 

Python端TFLite转换代码

converter = tf.lite.TFLiteConverter.from_saved_model('/content/GarbageClassificator_2Labels')
converter.optimizations = [tf.lite.Optimize.DEFAULT]
converter.target_spec.supported_ops = [tf.lite.OpsSet.TFLITE_BUILTINS, tf.lite.OpsSet.SELECT_TF_OPS]
converter.allow_custom_ops = True
converter.inference_input_type = tf.float32
converter.inference_output_type = tf.float32

tflite_model = converter.convert()

Python端TFLite模型执行代码

# 加载TFLITE模型
ruta_modelo_tflite = "garbageTflite_160x120.tflite"
interpreter = tf.lite.Interpreter(model_path=ruta_modelo_tflite)
interpreter.allocate_tensors()

input_details = interpreter.get_input_details()
output_details = interpreter.get_output_details()

print(input_details)
print(output_details)
predictions = []
for num in range(0, len(X_test)): # X_test已归一化(x_test / 255)
  input_data = np.array(X_test[num], dtype=np.float32)
  input_data = np.expand_dims(input_data, axis=0)

  interpreter.set_tensor(input_details[0]['index'], input_data)

  interpreter.invoke()

  output_data = interpreter.get_tensor(output_details[0]['index'])

  predictions.append(np.argmax(output_data))

输入输出详情

输入详情

[{'name': 'serving_default_input_1:0', 'index': 0, 'shape': array([  1, 160, 120,   3], dtype=int32), 'shape_signature': array([ -1, 160, 120,   3], dtype=int32), 'dtype': <class 'numpy.float32'>, 'quantization': (0.0, 0), 'quantization_parameters': {'scales': array([], dtype=float32), 'zero_points': array([], dtype=int32), 'quantized_dimension': 0}, 'sparsity_parameters': {}}]

输出详情

[{'name': 'StatefulPartitionedCall:0', 'index': 100, 'shape': array([1, 2], dtype=int32), 'shape_signature': array([-1,  2], dtype=int32), 'dtype': <class 'numpy.float32'>, 'quantization': (0.0, 0), 'quantization_parameters': {'scales': array([], dtype=float32), 'zero_points': array([], dtype=int32), 'quantized_dimension': 0}, 'sparsity_parameters': {}}]

排查与解决思路

1. 输入数据格式匹配问题

  • 图像格式转换:ESP32CAM默认输出格式可能是YUV或JPEG,而非模型要求的RGB888。直接读取fb->buf会导致输入数据完全错误,需要将摄像头输出转换为RGB格式后再输入模型。
  • 维度与通道顺序:Python端输入维度是[batch, height, width, channels](1,160,120,3),且通道为RGB。需确认ESP32CAM输出的图像维度顺序、通道顺序是否与训练时一致,若为BGR或其他顺序,需调整后再输入。

2. 模型算子支持问题

  • 算子注册完整性:MobileNetV2依赖ReLU6、Mul等算子,当前MicroMutableOpResolver未添加这些算子。可尝试替换为tflite::AllOpsResolver(注意内存占用),或通过TFLite模型分析工具查看模型依赖的所有算子,确保全部注册。
  • 转换参数优化:转换模型时使用了SELECT_TF_OPS,但TFLite Micro对部分TF原生算子支持有限。建议移除SELECT_TF_OPS和allow_custom_ops,仅使用TFLITE_BUILTINS重新转换模型。

3. 内存与张量问题

  • Arena内存大小:当前分配的1900000字节已接近使用上限,可尝试增大kArenaSize(如2000000),避免内存不足导致模型推理异常。
  • 输入张量验证:检查model_input_buffer是否正确指向输入张量的内存区域,可在输入数据填充后打印部分数值,确认与摄像头实际输出一致,避免野指针或内存越界。

4. 推理流程问题

  • 错误信息捕获:interpreter->Invoke()返回非kTfLiteOk时,仅打印了"Error",需通过error_reporter获取详细错误信息,定位推理失败的具体原因。
  • 输入重置:确保每次推理前输入张量被正确覆盖,避免残留数据影响结果。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.08 06:18:10