为何ImportExampleGen读取TFRecords返回SparseTensor而非Tensor?
问题背景与现象
我将CSV文件转换为TFRecords文件的操作如下:
源CSV文件:./dataset/csv/file.csv
feature_1, feture_2, output 1, 1, 1 2, 2, 2 3, 3, 3
转换为TFRecords的代码
import tensorflow as tf import csv import os print(tf.__version__) def create_csv_iterator(csv_file_path, skip_header): with tf.io.gfile.GFile(csv_file_path) as csv_file: reader = csv.reader(csv_file) if skip_header: # Skip the header next(reader) for row in reader: yield row def _int64_feature(value): """Returns an int64_list from a bool / enum / int / uint.""" return tf.train.Feature(int64_list=tf.train.Int64List(value=[value])) def create_example(row): """ Returns a tensorflow.Example Protocol Buffer object. """ features = {} for feature_index, feature_name in enumerate(["feature_1", "feture_2", "output"]): feature_value = row[feature_index] features[feature_name] = _int64_feature(int(feature_value)) return tf.train.Example(features=tf.train.Features(feature=features)) def create_tfrecords_file(input_csv_file): """ Creates a TFRecords file for the given input data """ output_tfrecord_file = input_csv_file.replace("csv", "tfrecords") writer = tf.io.TFRecordWriter(output_tfrecord_file) print("Creating TFRecords file at", output_tfrecord_file, "...") for i, row in enumerate(create_csv_iterator(input_csv_file, skip_header=True)): if len(row) == 0: continue example = create_example(row) content = example.SerializeToString() writer.write(content) writer.close() print("Finish Writing", output_tfrecord_file)
执行转换:
create_tfrecords_file("./dataset/csv/file.csv")
使用TFX读取TFRecords的流程
import os import absl import tensorflow_model_analysis as tfma tf.get_logger().propagate = False from tfx import v1 as tfx from tfx.orchestration.experimental.interactive.interactive_context import InteractiveContext %load_ext tfx.orchestration.experimental.interactive.notebook_extensions.skip
初始化上下文并读取数据:
context = InteractiveContext() example_gen = tfx.components.ImportExampleGen(input_base="./dataset/tfrecords") context.run(example_gen, enable_cache=True)
生成统计信息:
statistics_gen = tfx.components.StatisticsGen( examples=example_gen.outputs['examples']) context.run(statistics_gen, enable_cache=True)
生成Schema:
schema_gen = tfx.components.SchemaGen( statistics=statistics_gen.outputs['statistics'], infer_feature_shape=False) context.run(schema_gen, enable_cache=True)
Transform组件代码
文件:./transform.py
def preprocessing_fn(inputs): """tf.transform's callback function for preprocessing inputs. Args: inputs: map from feature keys to raw not-yet-transformed features. Returns: Map from string feature key to transformed feature operations. """ print(inputs) return inputs
运行Transform:
transform = tfx.components.Transform( examples=example_gen.outputs['examples'], schema=schema_gen.outputs['schema'], module_file=os.path.abspath("./transform.py")) context.run(transform, enable_cache=True)
问题
在preprocessing_fn函数中,我发现inputs是SparseTensor对象。我的数据集样本为密集型,本应返回Tensor,请问这是为何?我是否存在操作错误?
原因与解决方案
出现这个问题的核心原因是你在生成Schema时设置了infer_feature_shape=False:
schema_gen = tfx.components.SchemaGen( statistics=statistics_gen.outputs['statistics'], infer_feature_shape=False)
当infer_feature_shape=False时,TFX无法确定每个特征的固定形状,会默认将所有特征解析为SparseTensor。而你的数据集是密集型的,每个样本的特征都有固定的单值,应该让Schema明确特征的形状。
具体解决步骤
- 修改SchemaGen参数:移除
infer_feature_shape=False,让TFX自动推断特征的固定形状:
schema_gen = tfx.components.SchemaGen( statistics=statistics_gen.outputs['statistics']) context.run(schema_gen, enable_cache=True)
如果自动推断不符合预期,也可以手动编写Schema文件,明确指定每个特征的类型和形状。
- 清除缓存重新运行:因为之前启用了组件缓存,修改参数后需要禁用缓存重新运行,确保新的Schema生效:
context.run(schema_gen, enable_cache=False) context.run(transform, enable_cache=False)
- 验证TFRecords写入正确性:你的TFRecords写入代码是正确的,每个特征都被存储为单值int64类型,本身属于密集数据,只要Schema正确识别形状,Transform组件就会将其解析为Tensor而非SparseTensor。
内容的提问来源于stack exchange,提问作者Mehran
相关产品推荐
相关产品推荐

