如何用已训练多类语义分割模型预测测试库所有图像并保存输出?
问题原因及修复方案
核心问题排查点
- 确认Google Drive已正确挂载:
/content/drive路径需要先执行Drive挂载代码才能读写,未挂载的情况下写入该路径会静默失败。 - 确认保存目录存在:cv2.imwrite不会自动创建父目录,如果
/content/drive/MyDrive/BD_filtred文件夹不存在,会直接保存失败无报错,需先手动创建该文件夹,或者在代码里增加自动创建目录的逻辑。 - 数据类型不匹配:
np.argmax返回的数组是int64类型,cv2不支持该类型的图像写入,需要转换为uint8(类别数≤255时使用)或者uint16(类别数超过255时使用)类型才能正常保存。 - 代码逻辑冗余:原代码每次循环都对整个测试集执行一次预测,不仅运行效率极低,还可能因重复预测引发不必要的异常。
修复后的代码
import os import cv2 import numpy as np # 自动创建保存目录,不存在则新建,已存在也不会报错 save_dir = '/content/drive/MyDrive/BD_filtred/' os.makedirs(save_dir, exist_ok=True) # 一次性预测所有测试集,避免循环重复预测 y_pred = model.predict(test_images) # 转换为cv2支持的uint8类型,适配最多255类的分割场景 y_pred_argmax = np.argmax(y_pred, axis=3).astype(np.uint8) # 循环保存所有预测结果 for img_idx in range(test_images.shape[0]): save_path = os.path.join(save_dir, f'ok{img_idx+1}.png') cv2.imwrite(save_path, y_pred_argmax[img_idx]) # 打印日志确认保存进度 print(f"已保存预测结果:{save_path}")
额外排查方法
如果运行后还是没有文件生成,可以在cv2.imwrite后增加返回值判断,快速定位写入异常:
write_status = cv2.imwrite(save_path, y_pred_argmax[img_idx]) if not write_status: print(f"写入失败,对应路径:{save_path}")
内容的提问来源于stack exchange,提问作者Matança de Porco
相关产品推荐
相关产品推荐

