多轮成功训练后突发TensorFlow image.decode_png PNG数据损坏错误的问题求助
问题分析与解决方案
首先,你遇到的InvalidArgumentError: Invalid PNG data错误,在训练多轮后突然出现且文件训练前完好,大概率是磁盘I/O异常、文件系统问题,或者是TensorFlow图模式下异常处理失效导致的问题。下面一步步拆解原因和解决办法:
一、先排查硬件与文件系统层面的问题
这是最可能的根因,因为训练前文件完好,训练中损坏,说明磁盘在读写过程中出现了异常:
- 检查磁盘健康:用磁盘检测工具(比如Linux的
smartctl,Windows的磁盘检查工具)排查磁盘是否有坏道、I/O错误。如果是机械硬盘,坏道会导致文件读取/写入时损坏;如果是SSD,可能是寿命到期或固件问题。 - 验证文件完整性:重新对比训练前备份的PNG文件哈希值(比如MD5),找出所有损坏的文件,从备份中替换它们。不要直接继续使用损坏的文件,否则错误会反复出现。
- 检查磁盘空间:确保训练所在磁盘有足够剩余空间,空间不足会导致文件系统写入异常,甚至截断正在读取的文件。
二、修复代码中的异常处理逻辑
你的load_image方法里用了Python的try-except,但TensorFlow的tf.io.read_file和tf.image.decode_png是在图模式下执行的,Python的异常捕获无法拦截TensorFlow op的错误,导致错误直接抛出中断训练。需要调整代码,让异常处理在eager模式下生效:
修改load_image为使用tf.py_function包装
这样可以把Python的解码逻辑嵌入到tf.data流程中,同时正常捕获异常:
def load_image(self, filename, odometry): def py_load(filename1, filename2, odom): # 把Tensor转换为Python字符串路径 path1 = filename1.numpy().decode('utf-8') path2 = filename2.numpy().decode('utf-8') try: # 先读取文件数据 img1_data = tf.io.read_file(path1).numpy() img2_data = tf.io.read_file(path2).numpy() # 用PIL提前验证PNG有效性(比TensorFlow的检查更严格) from PIL import Image import io Image.open(io.BytesIO(img1_data)).verify() Image.open(io.BytesIO(img2_data)).verify() # 再用TensorFlow解码处理 img1 = self.decode_img(tf.convert_to_tensor(img1_data)) img2 = self.decode_img(tf.convert_to_tensor(img2_data)) img = tf.concat([img1, img2], -1) return img, odom except Exception as e: print(f"⚠️ Corrupted file detected: {path1} or {path2}") print(f"Error details: {str(e)}") # 返回一个无效样本标记,后续过滤掉 return tf.zeros((self.height, self.width, 6), dtype=tf.float32), tf.zeros((6,), dtype=tf.float32) # 用tf.py_function包装Python逻辑,指定输出类型 return tf.py_function( py_load, inp=[filename[0], filename[1], odometry], Tout=(tf.float32, tf.float32) )
添加无效样本过滤
在构建dataset后,过滤掉我们返回的无效标记样本:
# 在你的__init__方法中,dataset.batch(batch_size)之后添加 dataset = dataset.filter(lambda x, y: tf.reduce_sum(x) != 0)
三、提前校验所有训练数据
在训练开始前,先遍历所有PNG文件,提前找出损坏的文件,避免训练中途出错。可以在VisualOdometryDataLoader类中添加一个校验方法:
def validate_all_images(self, image_paths): from PIL import Image import io corrupted_files = [] for pair_paths in image_paths: for path in pair_paths: try: with open(path, 'rb') as f: img_bytes = f.read() Image.open(io.BytesIO(img_bytes)).verify() except Exception as e: corrupted_files.append(f"{path}: {str(e)}") if corrupted_files: print(f"❌ Found {len(corrupted_files)} corrupted images:") for err in corrupted_files: print(err) raise ValueError("Corrupted image files detected. Please fix or replace them before training.")
然后在__init__中调用这个方法:
def __init__(self, datapath, height, width, batch_size, test=False, val=False, sequence_test='10'): # ... 原有代码 ... images_stacked, odometries = self.get_data() # 提前校验所有图片 self.validate_all_images(images_stacked) dataset = tf.data.Dataset.from_tensor_slices((images_stacked, odometries)) # ... 原有代码 ...
四、额外优化建议
- 使用SSD存储训练数据:SSD的I/O稳定性比机械硬盘高很多,能减少磁盘读写错误的概率。
- 备份训练数据:定期备份数据集,避免文件损坏后无法恢复。
- 检查系统日志:查看系统日志(比如Linux的
/var/log/syslog,Windows的事件查看器),确认训练过程中是否有进程崩溃、磁盘I/O错误等记录,定位根本原因。
内容的提问来源于stack exchange,提问作者Essam Goda
相关产品推荐
相关产品推荐

