如何读写(n,3)二维数组的TFRecord文件?代码无输出求助
看起来你在TFRecord读写的过程中踩了两个关键的小坑,导致代码没报错但就是不出结果,我帮你梳理清楚问题和解决办法:
核心问题分析
- Feature解析定义不匹配:你写入时,
arry_x/arry_y/arry_z都是长度为n的一维浮点数组,但读取时却用了FixedLenFeature([], tf.float32)——这表示你告诉TensorFlow每个feature是单个浮点值,和实际写入的数据结构完全不搭,相当于解析的时候找错了数据格式。 - 队列线程没启动:你用的
tf.train.string_input_producer依赖后台线程来把文件名喂给读取器,但你没启动这些线程,导致Session一直卡在等待数据的状态,自然不会输出任何内容。
另外,你用的TFRecordReader和string_input_producer都是TF1.x里比较老旧的API,现在更推荐用tf.data.TFRecordDataset,我会同时给你两种解决方案:一种是修复你现有旧API的写法,另一种是更简洁的新API写法。
方案1:修复旧API的读取代码
第一步:修正Feature解析规则
因为写入时每个feature是长度为n的数组,所以读取时要对应定义成相同长度的FixedLenFeature(如果n是固定值);如果n不固定,就用VarLenFeature:
def get_tfrecord_feature(n): return{ # 把n替换成你数组的实际长度,比如你的array是(100,3),就写[100] 'arry_x': tf.compat.v1.io.FixedLenFeature([n], tf.float32), 'arry_y': tf.compat.v1.io.FixedLenFeature([n], tf.float32), 'arry_z': tf.compat.v1.io.FixedLenFeature([n], tf.float32) }
第二步:启动队列线程
在Session里必须启动队列协调器和线程,这样数据才能被正常读取:
import tensorflow as tf # 替换成你自己的参数 n = 100 # 你的array的行数 file_name = "your_data.tfrecord" def get_tfrecord_feature(n): return{ 'arry_x': tf.compat.v1.io.FixedLenFeature([n], tf.float32), 'arry_y': tf.compat.v1.io.FixedLenFeature([n], tf.float32), 'arry_z': tf.compat.v1.io.FixedLenFeature([n], tf.float32) } filenames = [file_name] file_name_queue = tf.train.string_input_producer(filenames) reader = tf.TFRecordReader() _, serialized_example = reader.read(file_name_queue) data = tf.compat.v1.io.parse_single_example(serialized_example, features=get_tfrecord_feature(n)) x = data['arry_x'] y = data['arry_y'] z = data['arry_z'] # 这里batch_size=1会得到形状为(1, n)的张量 x_batch, y_batch, z_batch = tf.train.batch([x, y, z], batch_size=1) with tf.compat.v1.Session() as sess: # 启动队列相关的线程 coord = tf.train.Coordinator() threads = tf.train.start_queue_runners(sess=sess, coord=coord) try: print("读取到的x数组:") print(sess.run(x_batch)) print("\n读取到的y数组:") print(sess.run(y_batch)) print("\n读取到的z数组:") print(sess.run(z_batch)) except tf.errors.OutOfRangeError: print("所有数据都读完啦") finally: # 用完记得停止线程 coord.request_stop() coord.join(threads)
方案2:用tf.data API读取(推荐)
TF1.13之后推出的tf.data.TFRecordDataset写法更直观,也更容易维护,代码如下:
import tensorflow as tf n = 100 # 替换成你数组的实际行数 file_name = "your_data.tfrecord" def parse_example(serialized_example): # 定义和写入时匹配的feature描述 feature_description = { 'arry_x': tf.io.FixedLenFeature([n], tf.float32), 'arry_y': tf.io.FixedLenFeature([n], tf.float32), 'arry_z': tf.io.FixedLenFeature([n], tf.float32), } return tf.io.parse_single_example(serialized_example, feature_description) # 构建数据集 pipeline dataset = tf.data.TFRecordDataset(file_name) dataset = dataset.map(parse_example) dataset = dataset.batch(1) # 设置批量大小 # 创建迭代器读取数据 iterator = tf.compat.v1.data.make_one_shot_iterator(dataset) next_batch = iterator.get_next() with tf.compat.v1.Session() as sess: try: while True: data = sess.run(next_batch) print("读取到的x数组:") print(data['arry_x']) print("\n读取到的y数组:") print(data['arry_y']) print("\n读取到的z数组:") print(data['arry_z']) except tf.errors.OutOfRangeError: print("所有数据读取完毕")
额外小建议:更合理的TFRecord写入方式
如果你的(n,3)数组是n个三维样本(比如n个点的坐标),其实更推荐把每个样本单独写成一个Example,这样读取时可以灵活设置batch size,也符合TFRecord的常规使用场景。修改写入代码如下:
def _float_feature(value): return tf.train.Feature(float_list=tf.train.FloatList(value=value)) # 遍历每个样本,逐个写入TFRecord with tf.compat.v1.python_io.TFRecordWriter(file_name) as writer: for sample in array: example = tf.train.Example( features=tf.train.Features( feature={ 'x': _float_feature([sample[0]]), 'y': _float_feature([sample[1]]), 'z': _float_feature([sample[2]]) } ) ) writer.write(example.SerializeToString())
对应的读取代码就可以用你最初写的FixedLenFeature([], tf.float32),配合batch获取批量样本,使用起来更灵活。
内容的提问来源于stack exchange,提问作者dlfksj
相关产品推荐
相关产品推荐

