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

TensorFlow Lite模型推理问题:输入设置与结果解析求助

问题求助:TFLite模型批量推理的输入处理与指标计算问题

已训练并保存了*.tflite格式的TensorFlow Lite模型,编写了可选择模型文件与图片文件夹、对图片批量执行推理的代码,但遇到以下问题:

  • 不确定interpreter.set_tensor的使用是否正确
  • 模型输出仅得到[[255]],结果不符合预期
  • 不知如何针对按标签分类存储的图片(如plant类存于folder/plant目录)计算损失/准确率指标

当前代码如下:

def testModel(self, testData):
        #Test any model on any dataset
        model = "**path to model file**"

        #Loading TFLite model and allocating tensors.
        interpreter = tf.lite.Interpreter(model_path=model)
        interpreter.allocate_tensors()

        # Get input and output tensors.
        input_details = interpreter.get_input_details()
        output_details = interpreter.get_output_details()

        rawImg = "**path to test images folder**"

        imgNameList = glob.glob(os.path.join(os.getcwd(), rawImg) + os.sep + '*') # gets list of image names in dir

        #creates dataset and dataloader from images
        testDataset = SalObjDataset(img_name_list = imgNameList,lbl_name_list = [], transform=transforms.Compose([RescaleT(224),ToTensorLab(flag=0)]))
        testDataloader = DataLoader(testDataset,batch_size=1,shuffle=False)
        
        #loops through dataloader (goes through each image file)
        for _, data in enumerate(testDataloader):
            inputImg = data['image']

            if torch.cuda.is_available():
                inputImg = Variable(inputImg.cuda())
            else:
                inputImg = Variable(inputImg)
            
            #rearranges dimensions in image file to match the expected input dimensions
            #also changes the type to uint8 as expected
            inputImg = tf.transpose(inputImg.cpu(), perm = [0,2,3,1])
            inputImg = tf.cast(inputImg, tf.uint8)

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

            interpreter.invoke()

            output_data = interpreter.get_tensor_details()
            print(output_data)

if __name__ == '__main__':
    #initialise object with the modelID of the model you want to test
    #pass the testing data folder name to testModel()
    #this is the folder where the model is
    modelID = "model_1"
    tester = ModelTrainer(modelID)
    #this is the folder where the testing images are
    tester.testModel("model_1/model_1/plant")

问题解决建议

1. 修正interpreter.set_tensor的使用

  • 核对输入要求:先打印input_details[0]查看模型期望的输入形状、数据类型(多数TFLite模型输入为float32,而非uint8)。如果模型要求归一化的float32输入,需将图片张量从0-255的uint8转换为0-1或-1到1的float32,而非直接转成uint8。
  • 简化张量处理:PyTorch中无需使用Variable,可直接处理张量;设备转换逻辑可简化为:
    inputImg = inputImg.cpu() if torch.cuda.is_available() else inputImg
    
  • 维度转换验证:确认tf.transpose后的形状和input_details[0]['shape']完全匹配(比如模型输入是[1,224,224,3],转换后的张量形状需一致)。

2. 正确获取模型输出

你当前使用interpreter.get_tensor_details()是获取所有张量的元数据,而非模型的推理结果。正确获取输出的方式是:

# 推理后获取输出张量
output_data = interpreter.get_tensor(output_details[0]['index'])
print(output_data)

同时打印output_details[0],确认输出的形状、数据类型是否符合任务预期(比如分类任务应为[1, num_classes]的概率分布)。

3. 计算损失与准确率指标

  • 加载真实标签:由于图片按标签分目录存储,需从图片路径中提取标签。例如:

    # 替换原imgNameList的生成逻辑,同时获取标签
    imgNameList = []
    labelList = []
    # 遍历所有子目录
    for label_dir in os.listdir(rawImg):
        label_path = os.path.join(rawImg, label_dir)
        if os.path.isdir(label_path):
            # 获取该目录下所有图片
            imgs = glob.glob(os.path.join(label_path, '*'))
            imgNameList.extend(imgs)
            # 子目录名作为标签,可根据实际情况转为数字索引
            labelList.extend([label_dir]*len(imgs))
    

    之后将labelList传入SalObjDataset,确保数据加载时能获取真实标签。

  • 计算准确率:在推理循环中,将模型输出转换为预测标签(比如分类任务取概率最大的索引),与真实标签对比统计:

    correct = 0
    total = 0
    for idx, data in enumerate(testDataloader):
        inputImg = data['image']
        true_label = data['label']  # 假设数据集返回真实标签
        # ... 输入处理与推理逻辑 ...
        output_data = interpreter.get_tensor(output_details[0]['index'])
        # 根据任务类型转换预测结果,比如分类任务取argmax
        pred_label = tf.argmax(output_data, axis=1).numpy()[0]
        # 对比真实标签与预测标签(需确保标签格式一致,比如都是数字索引)
        if pred_label == true_label:
            correct +=1
        total +=1
    accuracy = correct / total
    print(f"准确率: {accuracy:.4f}")
    
  • 计算损失:TFLite模型默认仅保留推理逻辑,若需损失值,可手动计算:

    # 示例:交叉熵损失,需确保真实标签为one-hot格式或类别索引
    loss_fn = tf.keras.losses.SparseCategoricalCrossentropy()
    loss = loss_fn(true_label, output_data).numpy()
    

    或在导出TFLite模型时,将损失层一并导出(需修改模型导出逻辑)。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.19 11:00:59