【问题标题】:Tensorflow 1.10+: passing epoch to estimator input_fn?Tensorflow 1.10+:将纪元传递给估计器 input_fn?
【发布时间】: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


    【解决方案1】:

    我不知道任何提供epoch 数字作为参数的选项。

    也就是说,根据定义,一个时期是输入函数的一个特征,因此我们应该能够处理输入函数内部的所有内容,而不是访问训练参数。所以我认为你可以通过一点点摆弄来实现你所需要的。

    例如,如果我有 2 个数据集:ds1 和 ds2,只要“纪元”数字不能被 100 整除,我就想使用 ds1,那么我可以通过执行以下操作创建一个新数据集:

    dataset = ds1.repeat(99).concatenate(ds2)
    

    由于默认情况下会延迟加载数据集,因此我无需担心内存影响(我不会将 100 倍的数据加载到内存中)。

    显然,这确实对数据集的大小有影响,因此您需要考虑评估操作/回调等之间的步骤策略,但这应该很容易调整。

    【讨论】:

    • 感谢您的回复。不幸的是,这在我的情况下并不奏效。我的输入函数有一些 required 在运行时发生的每个处理。因此数据被加载,每 n 个 epoch 对数据应用一些随机性,然后基于此修改数据,然后返回 (features, labels) 元组/dataset。这是必需的,因为预先计算所有这些成本太高。
    猜你喜欢
    • 2018-05-03
    • 2018-06-03
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 2018-05-06
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    相关资源
    最近更新 更多