You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何读写(n,3)二维数组的TFRecord文件?代码无输出求助

看起来你在TFRecord读写的过程中踩了两个关键的小坑,导致代码没报错但就是不出结果,我帮你梳理清楚问题和解决办法:

核心问题分析

  1. Feature解析定义不匹配:你写入时,arry_x/arry_y/arry_z都是长度为n的一维浮点数组,但读取时却用了FixedLenFeature([], tf.float32)——这表示你告诉TensorFlow每个feature是单个浮点值,和实际写入的数据结构完全不搭,相当于解析的时候找错了数据格式。
  2. 队列线程没启动:你用的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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.14 08:24:52