【问题标题】:How to get class indices from a quantized TFLite?如何从量化的 TFLite 中获取类索引?
【发布时间】:2021-01-19 18:25:13
【问题描述】:

我一直在用 TensorFlow 训练一个量化的 Mobilenet V2,但我不知道如何从中获取类索引。

我使用的是 TensorFlow 1.12

以下是我的输入和输出详细信息。

Input details [{'name': 'normalized_input_image_tensor', 'index': 260, 'shape': array([  1, 300, 300,   3], dtype=int32), 'shape_signature': array([  1, 300, 300,   3], dtype=int32), 'dtype': <class 'numpy.uint8'>, 'quantization': (0.0078125, 128), 'quantization_parameters': {'scales': array([0.0078125], dtype=float32), 'zero_points': array([128], dtype=int32), 'quantized_dimension': 0}, 'sparsity_parameters': {}}]
Output details [{'name': 'TFLite_Detection_PostProcess', 'index': 252, 'shape': array([ 1, 10,  4], dtype=int32), 'shape_signature': array([ 1, 10,  4], 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': 'TFLite_Detection_PostProcess:1', 'index': 253, 'shape': array([ 1, 10], dtype=int32), 'shape_signature': array([ 1, 10], 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': 'TFLite_Detection_PostProcess:2', 'index': 254, 'shape': array([ 1, 10], dtype=int32), 'shape_signature': array([ 1, 10], 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': 'TFLite_Detection_PostProcess:3', 'index': 255, 'shape': array([1], dtype=int32), 'shape_signature': array([1], 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': {}}]

我一直在尝试通过执行以下操作来获取类索引:

interpreter = tf.lite.Interpreter(model_path=PATH)
interpreter.allocate_tensors()
input_details = interpreter.get_input_details()
output_details = interpreter.get_output_details()
interpreter.set_tensor(input_details[0]['index'], input_data)
interpreter.invoke()
classes = interpreter.get_tensor(output_details[1]['index'])[0]

但是,类索引不正确。打印时,classes 看起来像这样:[0. 0. 0. 0. 0. 0. 0. 0. 0. 0.]。我的数据集中有超过 1 个类,所以这没有任何意义。

获取类索引的正确方法是什么?

【问题讨论】:

  • 你的任务是分类还是物体检测?
  • 对象检测。

标签: python tensorflow tensorflow-lite


【解决方案1】:

经过大量实验,事实证明这不是量化问题。我们在创建 .tflite 时使用了错误的 graph_def .pb 文件,因此它预测了不存在的类。

【讨论】:

    【解决方案2】:

    尝试使用:

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

    【讨论】:

    • 我的问题不是索引类张量,而是值的问题。 classes 看起来像这样:[0. 0. 0. 0. 0. 0. 0. 0. 0. 0.] 当我打印它时。
    • 看起来模型什么也没预测。你可以使用tf2吗? tf1 不再受支持。
    • 据我所知,我们正在使用 slim,tf2 尚不支持。
    猜你喜欢
    • 1970-01-01
    • 1970-01-01
    • 2022-01-10
    • 2013-07-20
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    相关资源
    最近更新 更多