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

为何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明确特征的形状。

具体解决步骤

  1. 修改SchemaGen参数:移除infer_feature_shape=False,让TFX自动推断特征的固定形状:
schema_gen = tfx.components.SchemaGen(
    statistics=statistics_gen.outputs['statistics'])
context.run(schema_gen, enable_cache=True)

如果自动推断不符合预期,也可以手动编写Schema文件,明确指定每个特征的类型和形状。

  1. 清除缓存重新运行:因为之前启用了组件缓存,修改参数后需要禁用缓存重新运行,确保新的Schema生效:
context.run(schema_gen, enable_cache=False)
context.run(transform, enable_cache=False)
  1. 验证TFRecords写入正确性:你的TFRecords写入代码是正确的,每个特征都被存储为单值int64类型,本身属于密集数据,只要Schema正确识别形状,Transform组件就会将其解析为Tensor而非SparseTensor。

内容的提问来源于stack exchange,提问作者Mehran

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.16 21:40:27