【发布时间】:2019-06-14 08:38:16
【问题描述】:
tf.estimatorinput_fn 的签名可能如下所示:
def input_fn(files:list, params:dict):
dataset = tf.data.TFRecordDataset(files)
.map(lambda record: parse_record_fn(record))
if params['mode'] == 'train':
# train specific things
# ...
这样的定义允许一个人随后构造他们所有的input_fns,如下所示:
train_fn = lambda: input_fn(files['training_set'], {**params, **{"mode": "train"}})
valid_fn = lambda: input_fn(files['validation_set'], {**params, **{"mode": "eval"}})
test_fn = lambda: input_fn(files['test_set'], {**params, **{"mode": "test"}})
train_spec = tf.estimator.TrainSpec(input_fn=train_fn, ...)
eval_spec = tf.estimator.EvalSpec(input_fn=valid_fn, ...)
我的问题是如何更改input_fn 签名以允许基于时代的变化。我知道这可能会带来瓶颈,但如果我能做这样的事情会很好:
def input_fn(...):
# see above
epoch = params["epoch"]
if epoch % 100 == 0:
# modify or make a new dataset
# ...
return dataset.make_one_shot_iterator().get_next()
关键是要确保input_fn 仍然兼容:
tf.estimator.train_and_evaluate(estimator, train_spec, eval_spec)
【问题讨论】:
标签: python tensorflow tensorflow-estimator