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

如何在TFX的preprocessing_fn中获取特征列表、Schema及统计信息

问题

我有如下简单的TFX管道:

import os
from tfx import v1 as tfx


_dataset_folder = './tfrecords/train/*'
_pipeline_data_folder = './pipeline_data'
_serving_model_dir = os.path.join(_pipeline_data_folder, 'serving_model')

example_gen = tfx.components.ImportExampleGen(input_base=_dataset_folder)
statistics_gen = tfx.components.StatisticsGen(examples=example_gen.outputs['examples'])
schema_gen = tfx.components.SchemaGen(
    statistics=statistics_gen.outputs['statistics'],
    infer_feature_shape=True)
example_validator = tfx.components.ExampleValidator(
    statistics=statistics_gen.outputs['statistics'],
    schema=schema_gen.outputs['schema'])

_transform_module_file = 'preprocessing_fn.py'
transform = tfx.components.Transform(
    examples=example_gen.outputs['examples'],
    schema=schema_gen.outputs['schema'],
    module_file=os.path.abspath(_transform_module_file),
    custom_config={'statistics_gen': statistics_gen.outputs['statistics'],
                   'schema_gen': schema_gen.outputs['schema']})

_trainer_module_file = 'run_fn.py'
trainer = tfx.components.Trainer(
    module_file=os.path.abspath(_trainer_module_file),
    examples=transform.outputs['transformed_examples'],
    transform_graph=transform.outputs['transform_graph'],
    schema=schema_gen.outputs['schema'],
    train_args=tfx.proto.TrainArgs(num_steps=10),
    eval_args=tfx.proto.EvalArgs(num_steps=6))


pusher = tfx.components.Pusher(
  model=trainer.outputs['model'],
  push_destination=tfx.proto.PushDestination(
    filesystem=tfx.proto.PushDestination.Filesystem(
        base_directory=_serving_model_dir)))

components = [
    example_gen,
    statistics_gen,
    schema_gen,
    example_validator,
    transform,
    trainer,
    pusher,
]

pipeline = tfx.dsl.Pipeline(
    pipeline_name='straightforward_pipeline',
    pipeline_root=_pipeline_data_folder,
    metadata_connection_config=tfx.orchestration.metadata.sqlite_metadata_connection_config(
        f'{_pipeline_data_folder}/metadata.db'),
    components=components)

tfx.orchestration.LocalDagRunner().run(pipeline)

我已经在Transform步骤的custom_config中传入了statistics_gen和schema_gen的输出,现在需要实现自动遍历数据集特征并做转换,需要获取:

  • 数据集中的特征列表(不硬编码)
  • 每个特征的类型(不硬编码)
  • 每个特征的统计属性(如最小值、最大值,不硬编码)

我的问题是:如何在preprocessing_fn.py中实现上述需求?

我知道如果能访问CSV数据集可以用以下方式获取统计信息,但这会重复计算,而管道中statistics_gen和schema_gen已经完成了这些工作,我想直接复用它们的输出:

import tensorflow_data_validation as tfdv

dataset_stats = tfdv.generate_statistics_from_csv(examples_file)
feature_1_stats = tfdv.get_feature_stats(dataset_stats.datasets[0],
                                         tfdv.FeaturePath(['feature_1']))
解决方案

要在preprocessing_fn中复用statistics_gen和schema_gen的输出,需通过custom_config传递的Artifact URI加载对应文件,再用TFDV工具解析,具体实现如下:

1. 加载并解析Schema

Schema以schema.pbtxt格式存储在schema_gen Artifact的URI路径下,加载后可遍历所有特征及其类型:

import tensorflow_data_validation as tfdv
import tensorflow_transform as tft

def preprocessing_fn(inputs, custom_config=None):
    # 从custom_config获取Schema的URI并加载
    schema_uri = custom_config['schema_gen'].uri
    schema = tfdv.load_schema_text(f"{schema_uri}/schema.pbtxt")
    
    # 遍历特征,收集名称和类型
    feature_details = {}
    for feature in schema.feature:
        feature_details[feature.name] = {
            'type': feature.type
        }

2. 加载并解析统计数据

统计数据以stats_tfrecord格式存储在statistics_gen Artifact的URI路径下,加载后可提取特征的统计属性:

# 从custom_config获取统计数据的URI并加载
    stats_uri = custom_config['statistics_gen'].uri
    dataset_stats = tfdv.load_statistics(f"{stats_uri}/stats_tfrecord")
    
    # 为每个特征补充统计属性
    for feature_name, details in feature_details.items():
        feature_stats = tfdv.get_feature_stats(
            dataset_stats.datasets[0], 
            tfdv.FeaturePath([feature_name])
        )
        # 针对数值型特征提取min/max
        if details['type'] in [tfdv.FeatureType.FLOAT, tfdv.FeatureType.INT]:
            details['min'] = feature_stats.num_stats.min
            details['max'] = feature_stats.num_stats.max

3. 动态生成特征转换逻辑

基于收集到的特征信息,可自动生成对应的转换逻辑,比如数值型特征归一化、字符串特征词汇表映射:

# 构建输出特征
    outputs = {}
    for feature_name, details in feature_details.items():
        if details['type'] in [tfdv.FeatureType.FLOAT, tfdv.FeatureType.INT]:
            # 数值型特征归一化到[0,1]
            outputs[feature_name] = tft.scale_to_0_1(inputs[feature_name])
        elif details['type'] == tfdv.FeatureType.STRING:
            # 字符串特征生成并应用词汇表
            outputs[feature_name] = tft.compute_and_apply_vocabulary(inputs[feature_name])
        # 可根据需求添加其他类型特征的处理逻辑
    
    return outputs

注意事项

  • 确保custom_config传递的Artifact URI在Transform组件运行时可访问,LocalDagRunner环境下默认无需额外配置,分布式环境需保证路径共享。
  • 处理嵌套特征时,需使用多层FeaturePath,例如tfdv.FeaturePath(['parent_feature', 'child_feature'])来定位嵌套子特征。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.06 16:25:34