训练自定义YOLO-NAS模型后,如何提取边界框坐标与标签?
提取YOLO-NAS预测结果的边界框与标签
针对你使用SuperGradients库的YOLO-NAS模型,以下是正确提取边界框(xyxy格式)和对应标签的方法:
import super_gradients.training as sgt # 加载预训练模型 best_model = sgt.models.get('yolo_nas_s', num_classes=1, checkpoint_path="C:/Users/Ritesh/Downloads/ckpt_best.pth") # 对单张图片执行预测 img = "cube.jpeg" predictions = best_model.predict(img) # 提取单张图片的预测结果(predict支持批量输入,此处取第一张图结果) single_pred = predictions[0] # 获取边界框(xyxy格式,顺序为x1, y1, x2, y2) bboxes = single_pred.bboxes_xyxy # 获取预测置信度 confidences = single_pred.confidence # 获取标签索引并转为整数 label_indices = single_pred.labels.astype(int) # 获取类别名称映射表 class_names = single_pred.class_names # 将标签索引映射为类别名称 pred_classes = [class_names[idx] for idx in label_indices] # 打印提取结果示例 for bbox, cls, conf in zip(bboxes, pred_classes, confidences): print(f"类别: {cls}, 置信度: {conf:.2f}, 边界框坐标: {bbox}")
关键说明
- 新版本SuperGradients中,
predict返回的Predictions对象支持直接通过索引访问单张图片的预测结果,无需调用私有属性_images_prediction_lst。 - 单张图片的预测结果对象直接提供
bboxes_xyxy、confidence、labels等公开属性,无需嵌套访问prediction.prediction层级。 - 若代码仍报错,建议升级SuperGradients到最新稳定版:
pip install --upgrade super-gradients
内容的提问来源于stack exchange,提问作者Ritesh Konka
相关产品推荐
相关产品推荐

