如何在Google Colab中加载并使用Kaggle垃圾分类数据集
加载并处理Kaggle垃圾分类数据集(12类)
问题背景
我已在Google Colab中完成垃圾分类数据集的解压,但无法通过代码加载并按12个类别处理。参考TensorFlow官方Fashion MNIST加载方式:
fashion_mnist = tf.keras.datasets.fashion_mnist (train_images, train_labels), (test_images, test_labels) = fashion_mnist.load_data()
尝试以下代码未成功运行,求正确实现方法:
import tensorflow as tf from tensorflow.keras.preprocessing.image import ImageDataGenerator # Define paths to your training and validation directories train_dir = 'garbage-classification/train' val_dir = 'garbage-classification/validation' # Create an ImageDataGenerator for data augmentation train_datagen = ImageDataGenerator(rescale=1./255) val_datagen = ImageDataGenerator(rescale=1./255) # Load images from directories train_generator = train_datagen.flow_from_directory( train_dir, target_size=(150, 150), # Resize images as needed batch_size=32, class_mode='categorical' # Use 'categorical' if you have multiple classes ) validation_generator = val_datagen.flow_from_directory( val_dir, target_size=(150, 150), batch_size=32, class_mode='categorical' )
修正方案与完整代码
1. 确认目录结构
flow_from_directory要求目录结构必须符合以下规则:
- 根目录(如
garbage-classification)下需包含train、validation子目录 - 每个子目录下,按12个类别创建单独文件夹(比如
train/cardboard、train/glass),文件夹内存放对应类别的图片
2. 修正Colab路径
Colab中解压的文件默认存放在/content/目录下,需补全路径确保指向正确:
train_dir = '/content/garbage-classification/train' val_dir = '/content/garbage-classification/validation'
可通过!ls /content命令查看当前目录下的实际文件夹名称,确保路径与实际文件夹名匹配。
3. 完整可运行代码
import tensorflow as tf from tensorflow.keras.preprocessing.image import ImageDataGenerator # 设定正确的数据集路径 train_dir = '/content/garbage-classification/train' val_dir = '/content/garbage-classification/validation' # 初始化数据生成器,可选添加数据增强提升模型泛化能力 train_datagen = ImageDataGenerator( rescale=1./255, rotation_range=20, width_shift_range=0.2, height_shift_range=0.2, horizontal_flip=True ) val_datagen = ImageDataGenerator(rescale=1./255) # 加载数据集 train_generator = train_datagen.flow_from_directory( train_dir, target_size=(150, 150), batch_size=32, class_mode='categorical', shuffle=True ) validation_generator = val_datagen.flow_from_directory( val_dir, target_size=(150, 150), batch_size=32, class_mode='categorical', shuffle=False ) # 验证加载结果 print(f"训练集类别数:{train_generator.num_classes}") print(f"训练集样本数:{train_generator.samples}") print(f"类别与索引映射:{train_generator.class_indices}")
4. 可选:转为Numpy数组格式(类似Fashion MNIST)
若需要得到和Fashion MNIST一致的Numpy数组格式,可使用image_dataset_from_directory实现:
# 加载为tf.data.Dataset对象 train_ds = tf.keras.utils.image_dataset_from_directory( train_dir, image_size=(150, 150), batch_size=32, label_mode='categorical' ) val_ds = tf.keras.utils.image_dataset_from_directory( val_dir, image_size=(150, 150), batch_size=32, label_mode='categorical' ) # 转换为Numpy数组 import numpy as np train_images = np.concatenate([x for x, y in train_ds], axis=0) train_labels = np.concatenate([y for x, y in train_ds], axis=0) val_images = np.concatenate([x for x, y in val_ds], axis=0) val_labels = np.concatenate([y for x, y in val_ds], axis=0)
内容的提问来源于stack exchange,提问作者Uais Amangeldi
相关产品推荐
相关产品推荐

