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
相关产品推荐
相关产品推荐

