You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.16 07:57:34