标准TensorFlow格式

TensorFlow的训练过程其实就是大量的数据在网络中不断流动的过程,而数据的来源在官方文档[^1](API r1.2)中介绍了三种方式,分别是:

  • Feeding。通过Python直接注入数据。
  • Reading from files。从文件读取数据,本文中的TFRecord属于此类方式。
  • Preloaded data。将数据以constant或者variable的方式直接存储在运算图中。

当数据量较大时,官方推荐采用标准TensorFlow格式[^2](Standard TensorFlow format)来存储训练与验证数据,该格式的后缀名为tfrecord。官方介绍如下:

A TFRecords file represents a sequence of (binary) strings. The format is not random access, so it is suitable for streaming large amounts of data but not suitable if fast sharding or other non-sequential access is desired.

从介绍不难看出,TFRecord文件适用于大量数据的顺序读取。而这正好是神经网络在训练过程中发生的事情。


如何使用TFRecord文件

对于TFRecord文件的使用,官方给出了两份示例代码,分别展示了如何生成与读取该格式的文件。

生成TFRecord文件

第一份代码convert_to_records.py [^3]将MNIST里的图像数据转换为了TFRecord格式 。仔细研读代码,可以发现TFRecord文件中的图像数据存储在Feature下的image_raw里。image_raw来自于data_set.images,而后者又来自mnist.read_data_sets()。因此images的真身藏在mnist.py这个文件里。

mnist.py并不难找,在Pycharm里按下ctrl后单击鼠标左键即可打开源代码。

继续追踪,可以在mnist里发现图像来自extract_images()函数。该函数的说明里清晰的写明:

Extract the images into a 4D uint8 numpy array [index, y, x, depth].
Args:
f: A file object that can be passed into a gzip reader.
Returns:
data: A 4D uint8 numpy array [index, y, x, depth].
Raises:
ValueError: If the bytestream does not start with 2051.

很明显,返回值变量名为data,是一个4D Numpy矩阵,存储值为uint8类型,即图像像素的灰度值(MNIST全部为灰度图像)。四个维度分别代表了:图像的个数,每个图像行数,每个图像列数,每个图像通道数。

在获得这个存储着像素灰度值的Numpy矩阵后,使用numpy的tostring()函数将其转换为Python bytes格式[^4],再使用tf.train.BytesList()函数封装为tf.train.BytesList类,名字为image_raw。最后使用tf.train.Example()image_raw和其它属性一遍打包,并调用tf.python_io.TFRecordWriter将其写入到文件中。

至此,TFRecord文件生成完毕。

可见,将自定义图像转换为TFRecord的过程本质上是将大量图像的像素灰度值转换为Python bytes,并与其它Feature组合在一起,最终拼接成一个文件的过程。

需要注意的是其它Feature的类型不一定必须是BytesList,还可以是Int64List或者FloatList。

读取TFRecord文件

第二份代码fully_connected_reader.py [1]展示了如何从TFRecord文件中读取数据。

读取数据的函数名为input()。函数内部首先通过tf.train.string_input_producer()函数读取TFRecord文件,并返回一个queue;然后使用read_and_decode()读取一份数据,函数内部用tf.decode_raw()解析出图像的灰度值,用tf.cast()解析出label的值。之后通过tf.train.shuffle_batch()的方法生成一批用来训练的数据。并最终返回可供训练的imageslabels,并送入inference部分进行计算。

在这个过程中,有以下几点需要留意:

  1. tf.decode_raw()解析出的数据是没有shape的,因此需要调用set_shape()函数来给出tensor的维度。
  2. read_and_decode()函数返回的是单个的数据,但是后边的tf.train.shuffle_batch()却能够生成批量数据。
  3. 如果需要对图像进行处理的话,需要放在第二项提到的两个函数中间。

其中第2点的原理我暂时没有弄懂。从代码上看read_and_decode()返回的是单个数据,shuffle_batch接收到的也是单个数据,不知道是如何生成批量数据的,猜测与queue有关系。

所以,读取TFRecord文件的本质,就是通过队列的方式依次将数据解码,并按需要进行数据随机化、图像随机化的过程。


参考


  1. Github: fully_connected_reader.py ↩︎

TFRecords转化和读取的更多相关文章

  1. [TFRecord格式数据]利用TFRecords存储与读取带标签的图片

    利用TFRecords存储与读取带标签的图片 原创文章,转载请注明出处~ 觉得有用的话,欢迎一起讨论相互学习~Follow Me TFRecords其实是一种二进制文件,虽然它不如其他格式好理解,但是 ...

  2. TensorFlow中数据读取之tfrecords

    关于Tensorflow读取数据,官网给出了三种方法: 供给数据(Feeding): 在TensorFlow程序运行的每一步, 让Python代码来供给数据. 从文件读取数据: 在TensorFlow ...

  3. 由浅入深之Tensorflow(3)----数据读取之TFRecords

    转载自http://blog.csdn.net/u012759136/article/details/52232266 原文作者github地址 概述 关于Tensorflow读取数据,官网给出了三种 ...

  4. tensorflowxun训练自己的数据集之从tfrecords读取数据

    当训练数据量较小时,采用直接读取文件的方式,当训练数据量非常大时,直接读取文件的方式太耗内存,这时应采用高效的读取方法,读取tfrecords文件,这其实是一种二进制文件.tensorflow为其内置 ...

  5. (第二章第四部分)TensorFlow框架之TFRecords数据的存储与读取

    系列博客链接: (第二章第一部分)TensorFlow框架之文件读取流程:https://www.cnblogs.com/kongweisi/p/11050302.html (第二章第二部分)Tens ...

  6. TensorFlow实践笔记(一):数据读取

    本文整理了TensorFlow中的数据读取方法,在TensorFlow中主要有三种方法读取数据: Feeding:由Python提供数据. Preloaded data:预加载数据. Reading ...

  7. Tensorflow高效读取数据

    关于Tensorflow读取数据,官网给出了三种方法: 供给数据(Feeding): 在TensorFlow程序运行的每一步, 让Python代码来供给数据. 从文件读取数据: 在TensorFlow ...

  8. VGGnet——从TFrecords制作到网络训练

    作为一个小白中的小白,多折腾总是有好处的,看了入门书和往上一些教程,很多TF的教程都是从MNIST数据集入手教小白入TF的大门,都是直接import MNIST,然后直接构建网络,定义loss和opt ...

  9. tensorflow学习笔记(10) mnist格式数据转换为TFrecords

    本程序 (1)mnist的图片转换成TFrecords格式 (2) 读取TFrecords格式 # coding:utf-8 # 将MNIST输入数据转化为TFRecord的格式 # http://b ...

随机推荐

  1. [Javascript Crocks] Apply a function in a Maybe context to Maybe inputs (curry & ap & liftA2)

    Functions are first class in JavaScript. This means we can treat a function like any other data. Thi ...

  2. vbs 脚本2

    一些很恶作剧的vbs程序代码 作者: 字体:[增加 减小] 类型:转载 时间:2013-01-16我要评论 恶作剧的vbs代码,这里提供的都是一些死循环或导致系统死机的vbs对机器没坏处,最多关机重启 ...

  3. Java-杂项:Float 加减精度问题

    ylbtech-Java-杂项:Float 加减精度问题 1.返回顶部 1. java float 加减精度问题在取这个字段的时候转换成BigDecimal就可以了同时,BigDecimal是可以设置 ...

  4. [jzoj 6073] 河 解题报告 (DP)

    interlinkage: https://jzoj.net/senior/#main/show/6073 description: solution: 考虑一条河$x$被染的效果 显然对于一条河$i ...

  5. SwiftUI 官方教程

    SwiftUI 官方教程 完整中文教程及代码请查看 https://github.com/WillieWangWei/SwiftUI-Tutorials   SwiftUI 官方教程 SwiftUI ...

  6. Python+unittest 接口自动化测试

    1.封装get.post#!/usr/bin/env python3# -*- coding: utf-8 -*- __author__ = 'hualai yu' import requests c ...

  7. C-字符串和格式化输入\输出

    1.字符串是一个或多个字符序列.字符串常量用双引号括起来“abc”,字符常量用单引号括起来‘’. 2.数组是同一类型的数据元素的有序序列.数据元素在内存中是连续存储的. C中没有为字符串定义专门的变量 ...

  8. javascript中天气接口案例

    <!DOCTYPE html> <html lang="en"> <head> <meta charset="UTF-8&quo ...

  9. VS 在代码中括号总是跟着类型后面

    if (OK.Text.Contains("运费")) {// 像这样子,这个大括号不是直接在IF下,而是跟在后面 工具-->选项-->文本编辑器-->C# -- ...

  10. Windows各种计时器

    (一):OnTimer类 1.打开对应对话框的类向导ClassWizard. 2.在消息映射MessageMaps中添加消息Message:WM_TIMER. 3.程序代码中将自动添加函数OnTime ...