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

Google Colab中导入带标签图片数据集:PyDrive操作卡点求助

解决Google Colab中带标签图片数据集的导入与读取问题

看起来你已经迈出了连接Google Drive的第一步,但后续的文件同步和数据集处理可以更顺畅。我会帮你梳理完整流程,包括优化Drive连接方式、同步数据到Colab,以及适配TensorFlow新版本的读取代码。

第一步:替换PyDrive,用原生方式挂载Google Drive

PyDrive的操作相对繁琐,Colab提供了更直观的原生挂载功能,能让你像访问本地文件一样操作Drive内容,步骤如下:

from google.colab import drive
drive.mount('/content/drive')

运行后会弹出授权链接,按照提示完成验证,你的Google Drive就会挂载到/content/drive/MyDrive/路径下。

假设你的data文件夹和Colab笔记本在同一个目录里,你可以先验证路径是否正确:

import os
# 替换成你的实际文件夹路径
data_dir = "/content/drive/MyDrive/Colab Notebooks/data/"
print(os.listdir(data_dir))
# 正常情况下会输出label1、label2这类文件夹名称

第二步:将Drive数据同步到Colab本地(可选但推荐)

直接从Drive读取文件速度较慢,建议把data文件夹复制到Colab的本地存储(/content/目录),训练时会更高效:

!cp -r "/content/drive/MyDrive/Colab Notebooks/data/" /content/
# 验证本地复制是否成功
print(os.listdir("/content/data/"))

第三步:用TensorFlow 2.x风格读取带标签数据集

你之前写的代码是TensorFlow 1.x的旧API,现在Colab默认使用TF 2.x,推荐用tf.data.Dataset或Keras内置工具来构建数据集,更简洁且兼容新版本。

方法1:用tf.keras.utils.image_dataset_from_directory(最简便)

Keras提供了直接从文件夹结构读取图片数据集的工具,无需手动处理路径和标签:

import tensorflow as tf
from tensorflow.keras.utils import image_dataset_from_directory

# 本地数据路径(如果直接用Drive读取就替换成上面的data_dir)
local_data_dir = "/content/data/"

# 构建训练集和验证集
batch_size = 32
img_height = 224
img_width = 224

train_ds = image_dataset_from_directory(
  local_data_dir,
  validation_split=0.2,  # 可选:按比例划分验证集
  subset="training",
  seed=123,
  image_size=(img_height, img_width),
  batch_size=batch_size)

val_ds = image_dataset_from_directory(
  local_data_dir,
  validation_split=0.2,
  subset="validation",
  seed=123,
  image_size=(img_height, img_width),
  batch_size=batch_size)

# 获取标签名称
class_names = train_ds.class_names
print("数据集标签列表:", class_names)

# 测试读取一批数据
for images, labels in train_ds.take(1):
  print("单张图片形状:", images[0].numpy().shape)
  print("对应标签:", labels[0].numpy(), "(对应类别:", class_names[labels[0].numpy()], ")")

方法2:手动用tf.data.Dataset构建(匹配你原来的思路)

如果你想手动处理路径和标签,这里修正你的旧代码以适配TF 2.x:

import tensorflow as tf
import glob

# 获取所有图片的路径
image_paths = glob.glob("/content/data/*/*.jpeg")

# 定义加载和预处理函数
def load_and_preprocess_image(path):
    # 读取并解码图片
    image = tf.io.read_file(path)
    image = tf.image.decode_jpeg(image, channels=3)
    # 调整图片尺寸
    image = tf.image.resize(image, [224, 224])
    # 归一化到0-1区间
    image = tf.cast(image, tf.float32) / 255.0
    
    # 从路径提取标签(路径格式:/content/data/labelX/img.jpeg)
    parts = tf.strings.split(path, "/")
    label = parts[-2]  # 倒数第二个元素是标签文件夹名
    # 如果标签是数字字符串,转换为整数类型
    label = tf.strings.to_number(label, out_type=tf.int32)
    return image, label

# 创建并处理数据集
dataset = tf.data.Dataset.from_tensor_slices(image_paths)
dataset = dataset.map(load_and_preprocess_image, num_parallel_calls=tf.data.AUTOTUNE)
# 打乱数据、分批、预加载提升效率
dataset = dataset.shuffle(buffer_size=len(image_paths)).batch(32).prefetch(tf.data.AUTOTUNE)

# 测试读取
for images, labels in dataset.take(1):
    print("批次图片形状:", images.shape)
    print("批次标签:", labels.numpy())

补充:你原来的PyDrive代码问题说明

你之前用drive.ListFile时的查询语句有误,正确的PyDrive获取文件夹内文件的方式是:

folder_id = "你的data文件夹ID"
file_list = drive.ListFile({'q': f"'{folder_id}' in parents and trashed=false"}).GetList()

但还是推荐用前面的原生挂载方式,操作更简单、读取速度更快。


内容的提问来源于stack exchange,提问作者Paul Schimmer

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.28 07:21:33