Tensorflow datasets.shuffle repeat batch方法
机器学习中数据读取是很重要的一个环节,TensorFlow也提供了很多实用的方法,为了避免以后时间久了又忘记,所以写下笔记以备日后查看。
1|0最普通的正常情况
首先我们看看最普通的情况:
# 创建0-10的数据集,每个batch取个数。dataset = tf.data.Dataset.range(10).batch(6)iterator = dataset.make_one_shot_iterator()next_element = iterator.get_next()with tf.Session() as sess: for i in range(2): value = sess.run(next_element) print(value)
输出结果
[0 1 2 3 4 5][6 7 8 9]
由结果我们可以知道TensorFlow能很好地帮我们自动处理最后一个batch的数据。
2|0datasets.batch(batch_size)与迭代次数的关系
但是如果上面for循环次数超过2会怎么样呢?也就是说如果 循环次数*批数量 > 数据集数量 会怎么样?我们试试看:
dataset = tf.data.Dataset.range(10).batch(6)iterator = dataset.make_one_shot_iterator()next_element = iterator.get_next()with tf.Session() as sess: >>==for i in range(3):==<< value = sess.run(next_element) print(value)
输出结果
[0 1 2 3 4 5][6 7 8 9]---------------------------------------------------------------------------OutOfRangeError Traceback (most recent call last)D:\Continuum\anaconda3\lib\site-packages\tensorflow\python\client\session.py in _do_call(self, fn, *args) 1277 try: ... ...省略若干信息... ... OutOfRangeError (see above for traceback): End of sequence [[Node: IteratorGetNext_64 = IteratorGetNext[output_shapes=[[?]], output_types=[DT_INT64], _device="/job:localhost/replica:0/task:0/device:CPU:0"](OneShotIterator_28)]]
可以知道超过范围了,所以报错了。
3|0datasets.repeat()
为了解决上述问题,repeat方法登场。还是直接看例子吧:
dataset = tf.data.Dataset.range(10).batch(6)dataset = dataset.repeat(2)iterator = dataset.make_one_shot_iterator()next_element = iterator.get_next()with tf.Session() as sess: for i in range(4): value = sess.run(next_element) print(value)
输出结果
[0 1 2 3 4 5][6 7 8 9][0 1 2 3 4 5][6 7 8 9]
可以知道repeat其实就是将数据集重复了指定次数,上面代码将数据集重复了2次,所以这次即使for循环次数是4也依旧能正常读取数据,并且都能完整把数据读取出来。同理,如果把for循环次数设置为大于4,那么也还是会报错,这么一来,我每次还得算repeat的次数,岂不是很心累?所以更简便的办法就是对repeat方法不设置重复次数,效果见如下:
dataset = tf.data.Dataset.range(10).batch(6)dataset = dataset.repeat()iterator = dataset.make_one_shot_iterator()next_element = iterator.get_next()with tf.Session() as sess: for i in range(6): value = sess.run(next_element) print(value)
输出结果:
[0 1 2 3 4 5][6 7 8 9][0 1 2 3 4 5][6 7 8 9][0 1 2 3 4 5][6 7 8 9]
此时无论for循环多少次都不怕啦~~
4|0datasets.shuffle(buffer_size)
仔细看可以知道上面所有输出结果都是有序的,这在机器学习中用来训练模型是浪费资源且没有意义的,所以我们需要将数据打乱,这样每批次训练的时候所用到的数据集是不一样的,这样啊可以提高模型训练效果。
另外shuffle前需要设置buffer_size:
不设置会报错,
buffer_size=1:不打乱顺序,既保持原序
buffer_size越大,打乱程度越大,演示效果见如下代码:
dataset = tf.data.Dataset.range(10).shuffle(2).batch(6)dataset = dataset.repeat(2)iterator = dataset.make_one_shot_iterator()next_element = iterator.get_next()with tf.Session() as sess: for i in range(4): value = sess.run(next_element) print(value)
输出结果:
[1 0 2 4 3 5][7 8 9 6][1 2 3 4 0 6][7 8 9 5]
注意:shuffle的顺序很重要,一般建议是最开始执行shuffle操作,因为如果是先执行batch操作的话,那么此时就只是对batch进行shuffle,而batch里面的数据顺序依旧是有序的,那么随机程度会减弱。不信你看:
dataset = tf.data.Dataset.range(10).batch(6).shuffle(10)dataset = dataset.repeat(2)iterator = dataset.make_one_shot_iterator()next_element = iterator.get_next()with tf.Session() as sess: for i in range(4): value = sess.run(next_element) print(value)
输出结果:
[0 1 2 3 4 5][6 7 8 9][0 1 2 3 4 5][6 7 8 9]
__EOF__
作 者:marsggbo
出 处:https://www.cnblogs.com/marsggbo
关于博主:热爱生活,爱读书/旅游,喜欢技术,乐于专研。评论和私信会在第一时间回复。或者直接私信我。
版权声明:署名 - 非商业性使用 - 禁止演绎,协议普通文本 | 协议法律文本。
出 处:https://www.cnblogs.com/marsggbo
关于博主:热爱生活,爱读书/旅游,喜欢技术,乐于专研。评论和私信会在第一时间回复。或者直接私信我。
版权声明:署名 - 非商业性使用 - 禁止演绎,协议普通文本 | 协议法律文本。
赞 (0)