训练TensorFlow Lite模型时如何丢弃数据集中会引发错误的元素
解决方案
方案1:正确适配tf.data.experimental.ignore_errors()
你遇到的AttributeError是因为ImageClassifierDataLoader是TensorFlow Lite Model Maker封装的高级数据对象,本身没有apply方法,你需要先访问它内部封装的原生tf.data.Dataset属性:
- 加载完训练集和验证集后,分别给两个数据集加上错误忽略逻辑:
# 给训练集添加错误静默丢弃逻辑 train_data.dataset = train_data.dataset.apply(tf.data.experimental.ignore_errors()) # 给验证集添加错误静默丢弃逻辑 validation_data.dataset = validation_data.dataset.apply(tf.data.experimental.ignore_errors()) # 之后正常执行训练代码即可 model = image_classifier.create(train_data, model_spec=model_spec.get('mobilenet_v2'), validation_data=validation_data)
如果使用的是TensorFlow 2.10及以上版本,也可以用更简洁的写法:
train_data.dataset = train_data.dataset.ignore_errors() validation_data.dataset = validation_data.dataset.ignore_errors()
注意:该方案会静默丢弃损坏的图片,建议先统计原始数据集和处理后数据集的样本数差值,确认丢弃的坏样本占比不会影响训练效果。
方案2:提前排查删除损坏图片(更推荐)
静默丢弃样本可能导致训练数据分布变化,更稳妥的方式是提前扫描整个数据集删除损坏的图片,参考代码如下:
import os import tensorflow as tf dataset_root = "你的数据集根目录路径" bad_img_list = [] # 遍历所有图片文件 for root, _, files in os.walk(dataset_root): for file in files: if file.lower().endswith(('.jpg', '.jpeg', '.png', '.bmp')): full_path = os.path.join(root, file) try: # 尝试读取解码图片,失败则标记为坏图 img_raw = tf.io.read_file(full_path) tf.io.decode_image(img_raw, channels=3, expand_animations=False) except Exception as e: bad_img_list.append(full_path) # 输出所有损坏图片路径 print(f"共发现{len(bad_img_list)}张损坏图片:") for img_path in bad_img_list: print(img_path) # 确认后可以取消注释自动删除坏图 # for img_path in bad_img_list: # os.remove(img_path)
内容的提问来源于stack exchange,提问作者user16571762
相关产品推荐
相关产品推荐

