【问题标题】:tf.data API cannot print all the batchestf.data API 无法打印所有批次
【发布时间】:2019-04-26 06:58:52
【问题描述】:

我正在自学tf.data API。我正在使用MNIST 数据集进行二进制分类。训练 x 和 y 数据一起压缩到完整的 train_dataset 中。与此 zip 方法链接在一起的首先是 batch() 数据集方法。数据以批量大小为 30 进行批量处理。由于我的训练集大小为 11623,批量大小为 128,因此我将有 91 个批量。最后一批的大小将是 103,这很好,因为这是 LSTM。此外,我正在使用辍学。当我计算批量准确度时,我关闭了 drop-out。

完整代码如下:

#Ignore the warnings
import warnings
warnings.filterwarnings("ignore")

import pandas as pd
import tensorflow as tf
import numpy as np

import matplotlib.pyplot as plt
plt.rcParams['figure.figsize'] = (8,7)

from tensorflow.examples.tutorials.mnist import input_data
mnist = input_data.read_data_sets("MNIST_data/")

Xtrain = mnist.train.images[mnist.train.labels < 2]
ytrain = mnist.train.labels[mnist.train.labels < 2]

print(Xtrain.shape)
print(ytrain.shape)

#Data parameters
num_inputs = 28
num_classes = 2
num_steps=28

# create the training dataset
Xtrain = tf.data.Dataset.from_tensor_slices(Xtrain).map(lambda x: tf.reshape(x,(num_steps, num_inputs)))
# apply a one-hot transformation to each label for use in the neural network
ytrain = tf.data.Dataset.from_tensor_slices(ytrain).map(lambda z: tf.one_hot(z, num_classes))
# zip the x and y training data together and batch and Prefetch data for faster consumption
train_dataset = tf.data.Dataset.zip((Xtrain, ytrain)).batch(128).prefetch(128)

iterator = tf.data.Iterator.from_structure(train_dataset.output_types,train_dataset.output_shapes)
X, y = iterator.get_next()

training_init_op = iterator.make_initializer(train_dataset)


#### model is here ####

#Network parameters
num_epochs = 2
batch_size = 128
output_keep_var = 0.5

with tf.Session() as sess:
    init.run()

    print("Initialized")
    # Training cycle
    for epoch in range(0, num_epochs):
        num_batch = 0
        print ("Epoch: ", epoch)
        avg_cost = 0.
        avg_accuracy =0
        total_batch = int(11623 / batch_size + 1)
        sess.run(training_init_op)
       while True:
            try:
                _, miniBatchCost = sess.run([trainer, loss], feed_dict={output_keep_prob: output_keep_var})
                miniBatchAccuracy = sess.run(accuracy, feed_dict={output_keep_prob: 1.0})
               print('Batch %d: loss = %.2f, acc = %.2f' % (num_batch, miniBatchCost, miniBatchAccuracy * 100))
                num_batch +=1
            except tf.errors.OutOfRangeError:
                break

当我运行这段代码时,它似乎正在工作并打印:

Batch 0: loss = 0.67276, acc = 0.94531
Batch 1: loss = 0.65672, acc = 0.92969
Batch 2: loss = 0.65927, acc = 0.89062
Batch 3: loss = 0.63996, acc = 0.99219
Batch 4: loss = 0.63693, acc = 0.99219
Batch 5: loss = 0.62714, acc = 0.9765
......
......
Batch 39: loss = 0.16812, acc = 0.98438
Batch 40: loss = 0.10677, acc = 0.96875
Batch 41: loss = 0.11704, acc = 0.99219
Batch 42: loss = 0.10592, acc = 0.98438
Batch 43: loss = 0.09682, acc = 0.97656
Batch 44: loss = 0.16449, acc = 1.00000

但是,正如我们很容易看到的那样,有些地方出了问题。只打印了 45 批而不是 91 批,我不知道为什么会这样。我尝试了很多东西,我想我错过了一些东西。

我可以使用repeat() 函数,但我不希望这样,因为我对最后一批有多余的观察,我希望 LSTM 来处理它。

【问题讨论】:

    标签: python tensorflow lstm tensorflow-datasets


    【解决方案1】:

    在直接基于tf.data 迭代器的get_next() 输出定义模型时,这是一个令人讨厌的陷阱。在您的循环中,您有两个 sess.run 调用,both 将迭代器前进一步。这意味着每个循环迭代实际上会消耗两个批次(而且您的损失和准确度计算是在不同批次上计算的)。

    不完全确定是否有解决此问题的“规范”方法,但您可以

    • 在与成本/训练步骤相同的run 调用中计算准确度。这意味着准确度计算也会受到 dropout 掩码的影响,但由于它是仅基于一个批次的近似值,所以这应该不是一个大问题。
    • 改为基于占位符定义您的模型,并在每次循环迭代中 run get_next 操作本身,然后将生成的 numpy 数组(即批处理)输入到损失/准确性计算中。

    【讨论】:

    • 是的,实际上我尝试了这两种方法。对于第二种方法,我特别不想使用占位符,因为我下一步将使用带有 SQL 查询的tf.data API,并且我想避免使用 feed_dict 来消费数据,因为这个 SQL 查询的结果很大。在我问了这个问题之后,我意识到我不能进行两次ses.run 调用,所以我将准确度放在与成本/培训步骤相同的run 调用中,因为这是你提到的第一种方法并且它有效。我认为这种方法更可行,我可以计算每个批次的精度,用于退出模型
    • 使用tf.data API 定义模型的其他方法是什么?它要么将tf.data 迭代器的输出分配给模型定义的每个参数,要么使用feed_dict 来使用sess.run([X,y])(或sess.run(next_batch))的输出,如herehere 中给出的那样。
    • 我不认为有任何其他“直接”方式,因为毕竟tf.data 迭代器只是返回张量,而get_next 操作将在每次调用时前进......一个选项可能是定义一个自定义数据集,在推进之前返回每个元素两次(或k 次,可以作为参数给出)。这可以例如通过数据集from_generator 实现。
    猜你喜欢
    • 1970-01-01
    • 1970-01-01
    • 2019-01-30
    • 2023-02-12
    • 1970-01-01
    • 2021-10-01
    • 2020-10-24
    • 1970-01-01
    • 2018-10-02
    相关资源
    最近更新 更多