TFX Pipeline运行正常但输出为空的原因排查
问题描述
我用以下代码实现了一个简单的TFX Pipeline,用于读取CSV文件并转换为TFRecord:
output_config = example_gen_pb2.Output(split_config= example_gen_pb2.SplitConfig(splits=[ example_gen_pb2.SplitConfig.Split(name='train', hash_buckets=8), example_gen_pb2.SplitConfig.Split(name='eval', hash_buckets=2) ]) ) example_gen = CsvExampleGen( input_base='data', output_config=output_config ) pipeline_root = 'artifacts' pipeline = Pipeline( pipeline_name='testing pipeline', pipeline_root=pipeline_root, components=[example_gen], enable_cache=True, metadata_connection_config=metadata.sqlite_metadata_connection_config( os.path.join('artifacts', 'metadata.sqlite') ) ) LocalDagRunner().run(pipeline)
我已经手动确认TFRecord文件生成正常,但打印Pipeline的输出字典时发现是空的,组件的输出也取不到内容:
print(pipeline.outputs) # 输出: {} print(example_gen.outputs['examples'].get()) # 输出: []
这个问题在.ipynb笔记本和.py脚本中都会出现,但使用InteractiveContext时没有这个问题,请问这是什么原因?
原因与解决方法
核心原因
pipeline.outputs为空的原因:Pipeline对象的outputs属性仅包含你显式声明为「Pipeline级输出」的组件输出。你在创建Pipeline时未指定outputs参数,所以它默认是空字典。组件
outputs.get()返回空的原因:LocalDagRunner运行Pipeline后,组件的输出元数据会被写入Metadata Store,但组件的outputs.get()方法不会自动从Metadata中读取已生成的结果路径。而InteractiveContext会在内部自动处理Metadata的查询和输出解析,所以不会出现这个问题。
解决方法
方法1:显式声明Pipeline级输出
如果你希望pipeline.outputs包含内容,可以在创建Pipeline时指定outputs参数:
pipeline = Pipeline( pipeline_name='testing pipeline', pipeline_root=pipeline_root, components=[example_gen], outputs={ 'examples': example_gen.outputs['examples'] }, enable_cache=True, metadata_connection_config=metadata.sqlite_metadata_connection_config( os.path.join('artifacts', 'metadata.sqlite') ) )
注意:即使声明了Pipeline级输出,pipeline.outputs['examples'].get()仍然不会直接返回实际路径,仍需通过Metadata Store查询。
方法2:从Metadata Store查询输出路径
要获取组件实际生成的TFRecord路径,需要手动连接Metadata Store查询:
from tfx.orchestration.metadata import Metadata # 连接到本地Metadata数据库 metadata_config = metadata.sqlite_metadata_connection_config( os.path.join('artifacts', 'metadata.sqlite') ) with Metadata(metadata_config) as store: # 获取最新的Pipeline运行记录 runs = store.get_runs() latest_run = runs[-1] # 查询该运行下的Examples类型产物 artifacts = store.get_artifacts_by_run_id(latest_run.id) for artifact in artifacts: if artifact.type_name == 'Examples': print(f"TFRecord输出路径: {artifact.uri}")
内容的提问来源于stack exchange,提问作者Sagnik Taraphdar
相关产品推荐
相关产品推荐

