使用Python基于JSON元数据整理Tensorflow图像分类数据集
适配TensorFlow图像分类任务的数据集整理方案
TensorFlow加载图像分类任务数据的标准目录结构为:数据集根目录下,每个子文件夹对应一个分类类别,子文件夹内存放该类别的所有图像,可直接通过tf.keras.utils.image_dataset_from_directory接口读取,无需额外编写复杂的数据加载逻辑。
实现逻辑
- 解析JSON元数据,提取每张图像的存储路径、对应类别标签
- 自动去重得到所有类别集合,批量创建类别对应的文件夹
- 按图像原始ID升序排序后,将图像迁移到对应类别目录,过程中自动处理文件缺失、重名冲突问题
- 支持复制/移动两种模式,避免误操作损坏原始数据
完整Python脚本
import json import os import shutil from pathlib import Path # -------------------------- 配置参数 按需修改 -------------------------- METADATA_JSON_PATH = "./metadata.json" # 元数据JSON文件路径 RAW_DATA_ROOT = "./" # 原始数据根目录(即JSON中image_filepath相对路径的根目录) OUTPUT_DATASET_ROOT = "./tf_cls_dataset" # 整理后数据集的输出根目录 COPY_INSTEAD_OF_MOVE = True # True=复制文件(保留原始数据) False=移动文件 # ---------------------------------------------------------------------- def main(): # 读取元数据 with open(METADATA_JSON_PATH, "r", encoding="utf-8") as f: metadata = json.load(f) # 创建输出根目录 Path(OUTPUT_DATASET_ROOT).mkdir(parents=True, exist_ok=True) # 收集所有类别,创建对应类别文件夹 all_classes = list({item["anomaly_class"] for item in metadata.values()}) for cls in all_classes: cls_dir = Path(OUTPUT_DATASET_ROOT) / cls cls_dir.mkdir(exist_ok=True) print(f"已创建共{len(all_classes)}个类别文件夹:{all_classes}") # 按图像ID升序排序,避免顺序错乱 sorted_items = sorted(metadata.items(), key=lambda x: int(x[0])) process_count = 0 miss_count = 0 for img_id, info in sorted_items: img_rel_path = info["image_filepath"] img_cls = info["anomaly_class"] src_path = Path(RAW_DATA_ROOT) / img_rel_path # 跳过不存在的源文件 if not src_path.exists(): print(f"[警告] 源文件不存在,跳过:{src_path}") miss_count +=1 continue # 构造目标路径,处理重名冲突 dst_dir = Path(OUTPUT_DATASET_ROOT) / img_cls dst_path = dst_dir / src_path.name rename_idx = 1 while dst_path.exists(): name_part, suffix = src_path.stem, src_path.suffix dst_path = dst_dir / f"{name_part}_{rename_idx}{suffix}" rename_idx +=1 # 执行复制/移动 if COPY_INSTEAD_OF_MOVE: shutil.copy2(src_path, dst_path) else: shutil.move(src_path, dst_path) process_count +=1 print(f"整理完成!共处理{process_count}张图像,跳过{miss_count}个缺失文件") print(f"数据集输出路径:{Path(OUTPUT_DATASET_ROOT).resolve()}") if __name__ == "__main__": main()
后续加载数据集方法
整理完成后,直接用TensorFlow内置接口即可加载标准化数据集,自动完成标签映射:
import tensorflow as tf # 加载训练集 train_ds = tf.keras.utils.image_dataset_from_directory( directory=OUTPUT_DATASET_ROOT, validation_split=0.2, # 按需设置验证集划分比例 subset="training", seed=123, image_size=(224, 224), # 按需设置模型要求的输入尺寸 batch_size=32 ) # 加载验证集 val_ds = tf.keras.utils.image_dataset_from_directory( directory=OUTPUT_DATASET_ROOT, validation_split=0.2, subset="validation", seed=123, image_size=(224, 224), batch_size=32 )
注意事项
- 首次运行建议保持
COPY_INSTEAD_OF_MOVE = True,确认整理结果无误后再切换为移动模式,避免原始数据丢失 - 若JSON中
image_filepath填写的是绝对路径,将RAW_DATA_ROOT设为空字符串即可 - 类别文件夹名称会直接作为模型训练时的类别标签名,若需要修改类别名称,直接重命名对应文件夹即可
内容的提问来源于stack exchange,提问作者Khairoo
相关产品推荐
相关产品推荐

