【问题标题】:How to split dataset into K-fold without loading the whole dataset at once?如何在不一次加载整个数据集的情况下将数据集拆分为 K 折叠?
【发布时间】:2021-07-07 12:58:52
【问题描述】:

我无法一次加载所有数据集,因此我使用tf.keras.preprocessing.image_dataset_from_directory() 在训练期间加载批量图像。如果我想将我的数据集分成 2 个子集(训练和验证),它工作得很好,但是,我想将我的数据集分成 K 折叠以进行交叉验证。 (5折就好了)

如何在不加载整个数据集的情况下制作 K 折叠? 我必须放弃使用tf.keras.preprocessing.image_dataset_from_directory() 吗?

【问题讨论】:

    标签: python tensorflow keras deep-learning k-fold


    【解决方案1】:

    我个人建议你切换到tf.data.Dataset()

    它不仅效率更高,而且在您可以实施的方面为您提供了更大的灵活性。

    假设你有图片(image_paths)和labels作为例子。

    这样,您可以创建如下管道:

    training_data = []
    validation_data = []
    kf = KFold(n_splits=5,shuffle=True,random_state=42)
    for train_index, val_index in kf.split(images,labels):
        X_train, X_val = images[train_index], images[val_index]
        y_train, y_val = labels[train_index], labels[val_index]
        training_data.append([X_train,y_train])
        validation_data.append([X_val,y_val])
    

    然后你可以创建类似的东西:

    for index, _ in enumerate(training_data):
        x_train, y_train = training_data[index][0], training_data[index][1]
        x_valid, y_valid = validation_data[index][0], validation_data[index][1]
       
        train_dataset = tf.data.Dataset.from_tensor_slices((x_train, y_train))
        train_dataset = train_dataset.map(mapping_function, num_parallel_calls=tf.data.experimental.AUTOTUNE)
        train_dataset = train_dataset.batch(batch_size)
        train_dataset = train_dataset.prefetch(buffer_size=tf.data.experimental.AUTOTUNE)
        
        
        validation_dataset = tf.data.Dataset.from_tensor_slices((x_valid, y_valid))
        validation_dataset = validation_dataset.map(mapping_function, num_parallel_calls=tf.data.experimental.AUTOTUNE)
        validation_dataset = validation_dataset.batch(batch_size)
        validation_dataset = validation_dataset.prefetch(buffer_size=tf.data.experimental.AUTOTUNE)
    
    
    
        model.fit(train_dataset,
                 validation_data=validation_dataset,
                 epochs=epochs,
                 verbose=2)
    

    【讨论】:

    • 感谢您的帮助。所以我应该把图像路径放在变量'images'中,然后将它们加载到mapping_function中?
    • 没错。映射函数作为逻辑应该有image = tf.io.decode_jpg(image_path,channels=3)等方法
    • 或您想要实现的其他处理:D
    猜你喜欢
    • 1970-01-01
    • 2018-08-09
    • 2011-10-28
    • 2020-09-27
    • 1970-01-01
    • 1970-01-01
    • 2016-11-01
    • 1970-01-01
    • 2020-02-03
    相关资源
    最近更新 更多