1.  打开下载链接:数据集-OpenDataLab      

下载安装包  cifar-100-python.tar.gz      

2. 用python解压文件(用7-zip直接解压会遇到问题)

import tarfile
import os

# 确保文件路径正确
tar_path = "D:/Dr/pycharm/Spiking-Wavelet-Transformer-main/cifar10-100/torch/cifar100/cifar-100-python.tar.gz"  # 源文件地址
extract_path = "D:/Dr/pycharm/Spiking-Wavelet-Transformer-main/cifar10-100/torch/cifar100"  # 解压后的地址

if not os.path.exists(extract_path):
    os.makedirs(extract_path)

with tarfile.open(tar_path, "r:gz") as tar:
    tar.extractall(path=extract_path)
print("解压完成!文件保存在:", os.path.abspath(extract_path))

3. 将解压后的文件提取为图片类型

import pickle
import numpy as np
from PIL import Image
import os


def unpickle(file):
    """读取 CIFAR-100 的二进制文件(train/test/meta)"""
    with open(file, 'rb') as fo:
        data = pickle.load(fo, encoding='bytes')
    return data


def save_images(data, output_dir, dataset_type, fine_label_names, coarse_label_names):
    """
    保存图片到指定文件夹,并按细分类名称创建子目录
    Args:
        data: 数据集(train 或 test 的二进制数据)
        output_dir: 输出根目录(如 'cifar-100-images')
        dataset_type: 'train' 或 'test'
        fine_label_names: 细分类名称列表(100类)
        coarse_label_names: 粗分类名称列表(20类)
    """
    images = data[b'data']  # 图像数据 (N, 3072)
    fine_labels = data[b'fine_labels']  # 细分类标签 (0-99)
    coarse_labels = data[b'coarse_labels']  # 粗分类标签 (0-19)

    # 创建数据集子目录(如 'cifar-100-images/train')
    save_dir = os.path.join(output_dir, dataset_type)
    os.makedirs(save_dir, exist_ok=True)

    # 遍历所有图片并保存
    for i in range(len(images)):
        # 转换图像格式 (3072 -> 32x32 RGB)
        img_flat = images[i]
        img_rgb = img_flat.reshape(3, 32, 32).transpose(1, 2, 0)  # (32, 32, 3)

        # 获取标签名称
        fine_label = fine_labels[i]
        coarse_label = coarse_labels[i]
        fine_name = fine_label_names[fine_label]
        coarse_name = coarse_label_names[coarse_label]

        # 创建细分类子目录(如 'cifar-100-images/train/apple')
        class_dir = os.path.join(save_dir, fine_name)
        os.makedirs(class_dir, exist_ok=True)

        # 保存图片(格式:序号_细分类_粗分类.png)
        img = Image.fromarray(img_rgb)
        img.save(os.path.join(class_dir, f'{i:05d}_{fine_name}_{coarse_name}.png'))

    print(f"{dataset_type} 图片已保存至 {save_dir}/")


def main():
    # 1. 设置输出目录
    output_dir = 'D:/Dr/pycharm/Spiking-Wavelet-Transformer-main/cifar10-100/torch/cifar100'
    os.makedirs(output_dir, exist_ok=True)

    # 2. 读取 meta 文件(获取类别名称)
    meta_data = unpickle('D:/Dr/pycharm/Spiking-Wavelet-Transformer-main/cifar10-100/torch/meta')
    fine_label_names = [name.decode('utf-8') for name in meta_data[b'fine_label_names']]  # 100类
    coarse_label_names = [name.decode('utf-8') for name in meta_data[b'coarse_label_names']]  # 20类

    # 3. 读取并保存训练集 (50,000 张)
    train_data = unpickle('D:/Dr/pycharm/Spiking-Wavelet-Transformer-main/cifar10-100/torch/train')
    save_images(train_data, output_dir, 'train', fine_label_names, coarse_label_names)

    # 4. 读取并保存测试集 (10,000 张)
    test_data = unpickle('D:/Dr/pycharm/Spiking-Wavelet-Transformer-main/cifar10-100/torch/test')
    save_images(test_data, output_dir, 'test', fine_label_names, coarse_label_names)

    print("所有图片保存完成!")


if __name__ == '__main__':
    main()

Logo

魔乐社区(Modelers.cn) 是一个中立、公益的人工智能社区,提供人工智能工具、模型、数据的托管、展示与应用协同服务,为人工智能开发及爱好者搭建开放的学习交流平台。社区通过理事会方式运作,由全产业链共同建设、共同运营、共同享有,推动国产AI生态繁荣发展。

更多推荐