【发布时间】: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