TensorFlow Object Detection API打印检测对象名称报错咨询
问题原因分析
你在处理detections结果时已经通过value[0, :num_detections].numpy()去掉了batch维度(输入为单张图片时的长度为1的第0维),所以:
classes = detections['detection_classes']是长度为num_detections的一维numpy数组,不是二维数组,classes[0]取到的是单个numpy.int64类型的数值,无法用于enumerate迭代,这就是报错的直接原因- 同理
scores = detections['detection_scores']也是一维数组,原来写法里的scores[0,index]也不存在对应的第二维度,后续也会报错 - 另外你取类别名称时需要和可视化逻辑一致,加上
label_id_offset=1的偏移,否则会出现类别id不匹配、取不到正确名称的问题
修复方案
直接替换报错的print行即可,以下是两种可用写法:
写法1(可读性更高)
threshold = 0.85 # 可以和可视化的min_score_thresh统一,避免打印结果和展示框不匹配 detected_names = [] for class_id, score in zip(classes, scores): if score > threshold: category_info = category_index.get(class_id + 1) # 加上类别偏移量 if category_info: detected_names.append(category_info['name']) print("检测到的对象:", detected_names)
写法2(简洁列表推导式)
print([category_index.get(class_id + 1, {}).get('name') for class_id, score in zip(classes, scores) if score > 0.8])
可选优化
如果需要同时获取检测对象的置信度、边界框坐标,可以同步遍历detections['detection_boxes']的结果,示例如下:
threshold = 0.85 detected_res = [] for class_id, score, box in zip(classes, scores, detections['detection_boxes']): if score > threshold: category_info = category_index.get(class_id + 1) if category_info: detected_res.append({ "名称": category_info['name'], "置信度": round(float(score), 4), "边界框": box.tolist() }) print(detected_res)
内容的提问来源于stack exchange,提问作者Sergio Ramirez Aguilar
相关产品推荐
相关产品推荐

