如何在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
相关产品推荐
相关产品推荐

