如何在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
相关产品推荐
相关产品推荐

