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

如何在Keras中通过flow_from_directory直接访问GCS中的图像?

用Keras ImageDataGenerator直接访问GCS中的图像目录

没问题,我来帮你搞定这个事儿~既然你已经在GCP上搭好GPU环境、把图像传到GCS了,接下来只需要几步就能让flow_from_directory直接读取GCS里的训练/验证/测试目录:

1. 先把必要的依赖装到位

如果你的Jupyter Notebook还没装GCS相关的依赖,先运行这行命令:

!pip install google-cloud-storage gcsfs

2. 确认GCS访问权限

如果你的Notebook是在GCP的Vertex AI或者Compute Engine上跑的,那默认已经有访问GCS的权限了,不用额外配置。要是你用的是其他环境,就得把服务账号密钥的路径设成环境变量:

import os
os.environ["GOOGLE_APPLICATION_CREDENTIALS"] = "/path/to/your/service-account-key.json"

3. 直接用GCS路径喂给flow_from_directory

TensorFlow 2.x及以后的版本里,tf.keras的flow_from_directory已经支持直接识别gs://开头的路径了,直接写就行!给你个完整的代码示例:

from tensorflow.keras.preprocessing.image import ImageDataGenerator

# 先初始化数据生成器,加你需要的数据增强参数
train_datagen = ImageDataGenerator(
    rescale=1./255,
    rotation_range=20,
    width_shift_range=0.2,
    height_shift_range=0.2,
    horizontal_flip=True
)

# 直接把GCS的训练目录路径传进去
train_generator = train_datagen.flow_from_directory(
    'gs://你的存储桶名称/train',  # 替换成你自己的GCS路径
    target_size=(224, 224),  # 改成你模型需要的图像尺寸
    batch_size=32,
    class_mode='categorical'  # 任务是分类的话,二分类就用'binary'
)

# 验证集和测试集照葫芦画瓢就行
val_datagen = ImageDataGenerator(rescale=1./255)
val_generator = val_datagen.flow_from_directory(
    'gs://你的存储桶名称/val',
    target_size=(224, 224),
    batch_size=32,
    class_mode='categorical'
)

test_datagen = ImageDataGenerator(rescale=1./255)
test_generator = test_datagen.flow_from_directory(
    'gs://你的存储桶名称/test',
    target_size=(224, 224),
    batch_size=32,
    class_mode='categorical',
    shuffle=False  # 测试集一般不用打乱顺序
)

要是上面的方法不好使?试试这个备选方案

如果你的TensorFlow版本比较旧,或者遇到了路径解析问题,可以用TensorFlow的Dataset API来读取GCS图像,再结合ImageDataGenerator的flow方法,性能还更好:

import tensorflow as tf
from tensorflow.keras.preprocessing.image import ImageDataGenerator

# 定义读取和预处理图像的函数
def load_image(path, label):
    img = tf.io.read_file(path)
    img = tf.image.decode_jpeg(img, channels=3)
    img = tf.image.resize(img, (224, 224))
    return img, label

# 从GCS目录创建数据集
train_dataset = tf.keras.utils.image_dataset_from_directory(
    'gs://你的存储桶名称/train',
    image_size=(224, 224),
    batch_size=32
)

# 初始化数据增强生成器
datagen = ImageDataGenerator(
    rotation_range=20,
    width_shift_range=0.2,
    height_shift_range=0.2,
    horizontal_flip=True
)

# 把tf数据集转成numpy数组,喂给flow方法
train_generator = datagen.flow(
    x=train_dataset.unbatch().map(lambda x, y: x).numpy(),
    y=train_dataset.unbatch().map(lambda x, y: y).numpy(),
    batch_size=32,
    shuffle=True
)

最后提几个注意点

  • 你的GCS目录结构必须符合flow_from_directory的要求:每个类别对应一个子目录,比如train/cat/、train/dog/这种,不然识别不了类别。
  • 要是遇到权限报错,去GCP控制台看看你的服务账号有没有GCS存储桶的读取权限,或者给存储桶加个Storage Object Viewer的IAM角色。
  • 数据集很大的话,优先用Dataset API的方案,它支持异步加载和预取,训练速度会更快。

内容的提问来源于stack exchange,提问作者Ryan Chase

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.25 07:16:58