如何制作可通过字符串键直接访问特征的TFRecord文件
问题原因
你报错的核心原因是跳过了TFRecord的读取、解析两个前置步骤,直接把存储文件路径的数据集传入了foo函数:
- 当前代码里的
example是文件路径的字符串张量,张量本身只支持整数、切片等类型的索引,用字符串作为键取值自然会触发类型错误。 - 他人代码中可以直接用
example["text"]取值,是因为他们在调用map(foo)之前,已经完成了TFRecord的解析,此时数据集的每个元素都是键为特征名、值为对应张量的字典,自然支持字符串下标取值。
正确实现步骤
你需要在原有逻辑前补全TFRecord读取和解析的逻辑,完整流程如下:
- 定义和你写入TFRecord时结构完全匹配的特征描述字典
- 编写解析函数,将二进制的TFRecord原始数据解析为特征字典
- 先读取TFRecord文件,再执行解析操作,最后才调用你的自定义处理函数
完整示例代码
import tensorflow as tf # 特征描述字典:必须和你写入TFRecord时的字段名、类型、维度完全一致 feature_description = { # 此处以text是字符串标量特征为例,你可以根据实际写入的特征调整 'text': tf.io.FixedLenFeature([], tf.string), # 其他存入的特征也需要在这里逐一声明 } # 解析函数:将二进制example解析为特征字典 def parse_tfrecord(example_proto): return tf.io.parse_single_example(example_proto, feature_description) def foo(example): # 此时example是解析后的字典,可直接用字符串键取值 text = example["text"] subtokens = some_other_function(text) features = { "my_subtokens": subtokens } return features input_files = ['test.tfrecord'] # 1. 读取TFRecord文件得到原始二进制数据集 d = tf.data.TFRecordDataset(input_files) # 2. 先执行解析,得到特征字典数据集 d = d.map(parse_tfrecord) # 3. 再调用你的自定义处理函数 d = d.map(foo)
注意事项
特征描述字典的配置必须和你写入TFRecord时的配置完全一致,否则会出现解析失败的问题。
内容的提问来源于stack exchange,提问作者Avery85
相关产品推荐
相关产品推荐

