Google Colab自定义链接导入数据集失败,请求技术支持
解决Google Colab中通过Drive链接导入数据集失败的问题
问题情况
在Google Colab中尝试通过自定义Google Drive下载链接导入数据集训练AI时,文件无法成功下载到Colab环境,执行代码触发FileNotFoundError,提示找不到目录/root/.keras/datasets/data_set。已尝试更换数据集文件夹、重新生成下载链接、修改路径等操作,问题仍未解决。
原代码
pip install tensorflow numpy matplotlib import tensorflow as tf from tensorflow.keras.models import Model from tensorflow.keras.applications import MobileNetV2, ResNet50, InceptionV3 # try to use them and see which is better from tensorflow.keras.layers import Dense from tensorflow.keras.callbacks import ModelCheckpoint, TensorBoard from tensorflow.keras.utils import get_file from tensorflow.keras.preprocessing.image import ImageDataGenerator import os import pathlib import numpy as np batch_size = 5 # 7 artists (currently) num_classes = 7 # training for 10 epochs epochs = 10 # size of each image IMAGE_SHAPE = (1000, 1000, 3) def load_data(): """This function downloads, extracts, loads, normalizes and one-hot encodes Flower Photos dataset""" # download the dataset and extract it #data_dir = get_file(origin='https://drive.google.com/uc?export=download&id=18FLNpct5RZlf4BZBpAaM6n7bzoQmSQR6', # original file https://storage.googleapis.com/download.tensorflow.org/example_images/flower_photos.tgz data_url = "https://drive.google.com/uc?export=download&id=18FLNpct5RZlf4BZBpAaM6n7bzoQmSQR6" data_dir = get_file(origin=data_url, fname="data_set", untar=True) # original fname = flower_photos data_dir = pathlib.Path(data_dir) # count how many images are there image_count = len(list(data_dir.glob('*/*.jpg'))) print("Number of images:", image_count) # get all classes for this dataset (types of flowers) excluding LICENSE file CLASS_NAMES = np.array([item.name for item in data_dir.glob('*') if item.name != "LICENSE.txt"]) # roses = list(data_dir.glob('roses/*')) # 20% validation set 80% training set image_generator = ImageDataGenerator(rescale=1/255, validation_split=0.2) # make the training dataset generator train_data_gen = image_generator.flow_from_directory(directory=str(data_dir), batch_size=batch_size, classes=list(CLASS_NAMES), target_size=(IMAGE_SHAPE[0], IMAGE_SHAPE[1]), shuffle=True, subset="training") # make the validation dataset generator test_data_gen = image_generator.flow_from_directory(directory=str(data_dir), batch_size=batch_size, classes=list(CLASS_NAMES), target_size=(IMAGE_SHAPE[0], IMAGE_SHAPE[1]), shuffle=True, subset="validation") return train_data_gen, test_data_gen, CLASS_NAMES def create_model(input_shape): # load MobileNetV2 model = MobileNetV2(input_shape=input_shape) # remove the last fully connected layer model.layers.pop() # freeze all the weights of the model except the last 4 layers for layer in model.layers[:-4]: layer.trainable = False # construct our own fully connected layer for classification output = Dense(num_classes, activation="softmax") # connect that dense layer to the model output = output(model.layers[-1].output) model = Model(inputs=model.inputs, outputs=output) # print the summary of the model architecture model.summary() # training the model using adam optimizer model.compile(loss="categorical_crossentropy", optimizer="adam", metrics=["accuracy"]) return model if __name__ == "__main__": # load the data generators train_generator, validation_generator, class_names = load_data() # constructs the model model = create_model(input_shape=IMAGE_SHAPE) # model name model_name = "MobileNetV2_finetune_last5" # some nice callbacks tensorboard = TensorBoard(log_dir=os.path.join("logs", model_name)) checkpoint = ModelCheckpoint(os.path.join("results", f"{model_name}" + "-loss-{val_loss:.2f}.h5"), save_best_only=True, verbose=1) # make sure results folder exist if not os.path.isdir("results"): os.mkdir("results") # count number of steps per epoch training_steps_per_epoch = np.ceil(train_generator.samples / batch_size) validation_steps_per_epoch = np.ceil(validation_generator.samples / batch_size) # train using the generators model.fit_generator(train_generator, steps_per_epoch=training_steps_per_epoch, validation_data=validation_generator, validation_steps=validation_steps_per_epoch, epochs=epochs, verbose=1, callbacks=[tensorboard, checkpoint])
报错信息
Downloading data from https://drive.google.com/uc?export=download&id=18FLNpct5RZlf4BZBpAaM6n7bzoQmSQR6 8192/Unknown - 0s 0us/stepNumber of images: 0 --------------------------------------------------------------------------- FileNotFoundError Traceback (most recent call last) <ipython-input-10-102c652a5bdf> in <cell line: 1>() 1 if __name__ == "__main__": 2 # load the data generators ----> 3 train_generator, validation_generator, class_names = load_data() 4 # constructs the model 5 model = create_model(input_shape=IMAGE_SHAPE) /usr/local/lib/python3.10/dist-packages/keras/src/preprocessing/image.py in __init__(self, directory, image_data_generator, target_size, color_mode, classes, class_mode, batch_size, shuffle, seed, data_format, save_to_dir, save_prefix, save_format, follow_links, subset, interpolation, keep_aspect_ratio, dtype) 561 if not classes: 562 classes = [] ---> 563 for subdir in sorted(os.listdir(directory)): 564 if os.path.isdir(os.path.join(directory, subdir)): 565 classes.append(subdir) FileNotFoundError: [Errno 2] No such file or directory: '/root/.keras/datasets/data_set'
问题分析
从报错日志可见,数据集仅下载了8192字节(远小于正常数据集大小),说明tensorflow.keras.utils.get_file无法正确处理Google Drive的下载链接——Google Drive对非公开或大文件会有验证机制,直接用uc?export=download链接经常导致下载不完整或失败,最终解压后的目录不存在或为空,触发后续的FileNotFoundError。
解决步骤
1. 确保Drive文件共享权限正确
- 打开Google Drive中的数据集文件,右键选择「获取共享链接」
- 设置权限为「知道链接的任何人都可查看」,保存后复制链接
2. 改用gdown库下载数据集(更可靠)
gdown是专门处理Google Drive文件下载的工具,能自动处理验证环节,替换原代码中的下载逻辑:
第一步:安装依赖
在Colab单元格中执行:
!pip install gdown tensorflow numpy matplotlib
第二步:修改load_data函数
将原load_data函数替换为以下代码:
def load_data(): """下载、提取、加载数据集并做预处理""" import gdown import tarfile # 替换为你的Drive文件ID(从共享链接中提取,比如链接里的18FLNpct5RZlf4BZBpAaM6n7bzoQmSQR6) file_id = "18FLNpct5RZlf4BZBpAaM6n7bzoQmSQR6" output_tar = "data_set.tgz" save_dir = "/root/.keras/datasets/" # 创建保存目录(如果不存在) os.makedirs(save_dir, exist_ok=True) # 下载数据集 gdown.download(f"https://drive.google.com/uc?id={file_id}", output_tar, quiet=False) # 解压到指定目录 with tarfile.open(output_tar, 'r:gz') as tar: tar.extractall(path=save_dir) # 定义数据集目录路径 data_dir = pathlib.Path(os.path.join(save_dir, "data_set")) # 验证目录是否存在 if not data_dir.exists(): raise FileNotFoundError(f"解压后数据集目录不存在,请检查文件是否为正确的压缩包:{data_dir}") # 统计图片数量 image_count = len(list(data_dir.glob('*/*.jpg'))) print("Number of images:", image_count) # 获取类别名称 CLASS_NAMES = np.array([item.name for item in data_dir.glob('*') if item.name != "LICENSE.txt"]) # 生成数据生成器 image_generator = ImageDataGenerator(rescale=1/255, validation_split=0.2) train_data_gen = image_generator.flow_from_directory( directory=str(data_dir), batch_size=batch_size, classes=list(CLASS_NAMES), target_size=(IMAGE_SHAPE[0], IMAGE_SHAPE[1]), shuffle=True, subset="training" ) test_data_gen = image_generator.flow_from_directory( directory=str(data_dir), batch_size=batch_size, classes=list(CLASS_NAMES), target_size=(IMAGE_SHAPE[0], IMAGE_SHAPE[1]), shuffle=True, subset="validation" ) return train_data_gen, test_data_gen, CLASS_NAMES
3. 验证数据集是否正确加载
修改代码后运行,先检查输出中的Number of images是否大于0,如果显示正常,说明数据集已正确下载并解压。
修改后的完整代码
!pip install gdown tensorflow numpy matplotlib import tensorflow as tf from tensorflow.keras.models import Model from tensorflow.keras.applications import MobileNetV2, ResNet50, InceptionV3 from tensorflow.keras.layers import Dense from tensorflow.keras.callbacks import ModelCheckpoint, TensorBoard from tensorflow.keras.preprocessing.image import ImageDataGenerator import os import pathlib import numpy as np batch_size = 5 num_classes = 7 epochs = 10 IMAGE_SHAPE = (1000, 1000, 3) def load_data(): """下载、提取、加载数据集并做预处理""" import gdown import tarfile # 替换为你的Drive文件ID file_id = "18FLNpct5RZlf4BZBpAaM6n7bzoQmSQR6" output_tar = "data_set.tgz" save_dir = "/root/.keras/datasets/" os.makedirs(save_dir, exist_ok=True) gdown.download(f"https://drive.google.com/uc?id={file_id}", output_tar, quiet=False) with tarfile.open(output_tar, 'r:gz') as tar: tar.extractall(path=save_dir) data_dir = pathlib.Path(os.path.join(save_dir, "data_set")) if not data_dir.exists(): raise FileNotFoundError(f"解压后数据集目录不存在,请检查文件是否为正确的压缩包:{data_dir}") image_count = len(list(data_dir.glob('*/*.jpg'))) print("Number of images:", image_count) CLASS_NAMES = np.array([item.name for item in data_dir.glob('*') if item.name != "LICENSE.txt"]) image_generator = ImageDataGenerator(rescale=1/255, validation_split=0.2) train_data_gen = image_generator.flow_from_directory( directory=str(data_dir), batch_size=batch_size, classes=list(CLASS_NAMES), target_size=(IMAGE_SHAPE[0], IMAGE_SHAPE[1]), shuffle=True, subset="training" ) test_data_gen = image_generator.flow_from_directory( directory=str(data_dir), batch_size=batch_size, classes=list(CLASS_NAMES), target_size=(IMAGE_SHAPE[0], IMAGE_SHAPE[1]), shuffle=True, subset="validation" ) return train_data_gen, test_data_gen, CLASS_NAMES def create_model(input_shape): model = MobileNetV2(input_shape=input_shape) model.layers.pop() for layer in model.layers[:-4]: layer.trainable = False output = Dense(num_classes, activation="softmax")(model.layers[-1].output) model = Model(inputs=model.inputs, outputs=output) model.summary() model.compile(loss="categorical_crossentropy", optimizer="adam", metrics=["accuracy"]) return model if __name__ == "__main__": train_generator, validation_generator, class_names = load_data() model = create_model(input_shape=IMAGE_SHAPE) model_name = "MobileNetV2_finetune_last5" tensorboard = TensorBoard(log_dir=os.path.join("logs", model_name)) checkpoint = ModelCheckpoint(os.path.join("results", f"{model_name}-loss-{val_loss:.2f}.h5"), save_best_only=True, verbose=1) os.makedirs("results", exist_ok=True) training_steps_per_epoch = np.ceil(train_generator.samples / batch_size) validation_steps_per_epoch = np.ceil(validation_generator.samples / batch_size) model.fit( train_generator, steps_per_epoch=training_steps_per_epoch, validation_data=validation_generator, validation_steps=validation_steps_per_epoch, epochs=epochs, verbose=1, callbacks=[tensorboard, checkpoint] )
注:已替换fit_generator为fit(TensorFlow 2.x中fit_generator已被弃用,fit直接支持生成器)
内容的提问来源于stack exchange,提问作者Wallah Billah
相关产品推荐
相关产品推荐

