【问题标题】:How to split dataset and feed into input_fn如何拆分数据集并输入 input_fn
【发布时间】:2021-05-13 13:37:06
【问题描述】:

我已经使用 input_fn() 中的 make_csv_dataset 成功读取了两个 csv 文件,并将其传递到 tf.estimator

我首先将主 csv 分成两个单独的帧,一个用于训练,一个用于测试,然后将它们保存为新的 csv 文件。

train, test = train_test_split(df, test_size = 0.2)

train_csv_path = 'data/2020_train.csv.gz'
test_csv_path = 'data/2020_test.csv.gz'

train.to_csv(train_csv_path, compression = 'gzip')
test.to_csv(test_csv_path, compression = 'gzip')
def make_input_fn(csv_path, n_epochs = None):
    def input_fn():
        dataset = tf.data.experimental.make_csv_dataset(csv_path,
                                                        batch_size = 1000,
                                                        label_name = 'Shipped On SSD',
                                                        compression_type = 'GZIP',
                                                        num_epochs = n_epochs)
        return dataset
    return input_fn

train_input_fn = make_input_fn(train_csv_path)
test_input_fn = make_input_fn(test_csv_path, n_epochs = 1)

但是,我只想使用一个文件并在数据集上进行拆分。

我可以成功拆分数据集(如this),但将其传递给tf.estimator 时会出现问题。我不知道如何使用在input_fn() 之外定义的数据集或如何在input_fn() 内进行拆分。

dataset = tf.data.experimental.make_csv_dataset(full_csv_path,
                                                batch_size = 1000,
                                                label_name = 'Shipped On SSD',
                                                compression_type = 'GZIP')

split = 4
split_fn = lambda *ds: ds[0] if len(ds) == 1 else tf.data.Dataset.zip(ds)

dataset_train = dataset.window(split, split + 1).flat_map(split_fn)
dataset_test = dataset.skip(split).window(1, split + 1).flat_map(split_fn)

【问题讨论】:

    标签: python tensorflow


    【解决方案1】:

    您可以将数据集创建包装在一个函数中。不幸的是,该函数将读取 csv 两次,每组一次。

    def make_input_fn_from_ds(ds, training=True):
      def input_fn():
        dataset = tf.data.experimental.make_csv_dataset(full_csv_path,
                                                    batch_size = 1000,
                                                    label_name = 'Shipped On SSD',
                                                    compression_type = 'GZIP')
    
        split = 4
        split_fn = lambda *ds: ds[0] if len(ds) == 1 else   tf.data.Dataset.zip(ds)
    
        if training:
          return dataset_train = dataset.window(split, split + 1).flat_map(split_fn)
        return dataset_test = dataset.skip(split).window(1, split + 1).flat_map(split_fn)
      return input_fn
    
    train_input_fn = make_input_fn_from_ds(dataset_train, training=True)
    test_input_fn = make_input_fn_from_ds(dataset_train, training=False)
    

    【讨论】:

    • 这给出了以下错误:迭代器的图()与图() 数据集:tf.Tensor(, shape=(), dtype=variant) 是在其中创建的。如果您使用的是 Estimator API,请确保 @ 返回的数据集没有任何部分987654322@ 函数在input_fn 函数之外定义。请确保管道中的所有数据集都创建在与迭代器相同的图中。
    • 哦,我忘记了这个限制。让我编辑我的答案。
    • 谢谢!我是否应该关闭洗牌,以免冒将训练数据集包含到测试数据集中的风险?
    • 是的,应该避免洗牌。或者您可以在随机播放之前调用tf.random.set_seed 使其具有确定性,并将reshuffle_each_iteration=False 传递给shuffle 方法。
    猜你喜欢
    • 1970-01-01
    • 2018-12-10
    • 1970-01-01
    • 1970-01-01
    • 2015-11-09
    • 1970-01-01
    • 2018-11-29
    • 1970-01-01
    • 1970-01-01
    相关资源
    最近更新 更多