多并发场景下如何获取当前Pipeline执行生成的Artifact URI?
TFX Pipeline运行后精准获取Artifact方案
问题背景
在TFX Pipeline执行完成后,需要对本次运行生成的Artifact进行后续操作,但现有查询方式存在局限性:
- 基础ML Metadata查询无法区分不同Pipeline运行实例;
- 依赖Pipeline名称的方案在同名Pipeline并发执行时会失效,无法定位到具体某次运行的Artifact。
核心需求:无论不同名称或同名Pipeline并发执行,都能准确获取刚完成的Pipeline运行所生成的Artifact。
示例TFX Pipeline代码
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) transform = tfx.components.Transform( examples=example_gen.outputs['examples'], schema=schema_gen.outputs['schema'], module_file=os.path.abspath('preprocessing_fn.py')) _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, transform, trainer, pusher, ] _pipeline_data_folder = './simple_pipeline_data' pipeline = tfx.dsl.Pipeline( pipeline_name='simple_pipeline', pipeline_root=_pipeline_data_folder, metadata_connection_config=tfx.orchestration.metadata.sqlite_metadata_connection_config( f'{_pipeline_data_folder}/metadata.db'), components=components) # 执行Pipeline并捕获Run ID runner = tfx.orchestration.LocalDagRunner() run_result = runner.run(pipeline) run_id = run_result.run_id
基础ML Metadata查询方式(无法区分运行实例)
这种方式只能查询所有Artifact,无法定位到某次Pipeline运行的产物:
import ml_metadata as mlmd connection_config = pipeline.metadata_connection_config store = mlmd.MetadataStore(connection_config) print(store.get_artifact_types())
现有部分解决方案(同名Pipeline并发失效)
以下函数通过Pipeline名称+组件名称查询最新Artifact,但无法区分同名Pipeline的不同运行实例:
def get_latest_artifact(metadata_connection_config, pipeline_name: str, component_name: str, type_name: str): with Metadata(metadata_connection_config) as metadata: context = metadata.store.get_context_by_type_and_name('node', f'{pipeline_name}.{component_name}') artifacts = metadata.store.get_artifacts_by_context(context.id) artifact_type = metadata.store.get_artifact_type(type_name) latest_artifact = max([a for a in artifacts if a.type_id == artifact_type.id], key=lambda a: a.last_update_time_since_epoch) artifact = types.Artifact(artifact_type) artifact.set_mlmd_artifact(latest_artifact) return artifact sqlite_path = './pipeline_data/metadata.db' metadata_connection_config = tfx.orchestration.metadata.sqlite_metadata_connection_config(sqlite_path) examples_artifact = get_latest_artifact(metadata_connection_config, 'simple_pipeline', 'SchemaGen', 'Schema')
可靠解决方案:基于Run ID精准定位
MLMD原生支持通过Run ID区分同一Pipeline的不同运行实例,每个Pipeline执行都会生成一个全局唯一的Run Context(类型为run)。以下是实现方案:
1. 获取本次运行的Run ID
在启动Pipeline时,从Orchestrator返回结果中提取Run ID:
- LocalDagRunner:
run()方法返回的结果包含run_id字段; - Kubeflow Pipelines:提交Pipeline后返回的Run对象包含唯一ID;
- Airflow:DAG运行的Execution ID可作为Run ID使用。
2. 基于Run ID查询Artifact的函数
import ml_metadata as mlmd from tfx.types import Artifact def get_artifacts_by_run_id(metadata_connection_config, run_id: str, component_name: str, type_name: str): with mlmd.MetadataStore(metadata_connection_config) as store: # 获取本次运行对应的Run Context run_context = store.get_context_by_type_and_name('run', run_id) if not run_context: raise ValueError(f"未找到Run ID: {run_id}") # 获取该Run下所有组件的Execution实例 executions = store.get_executions_by_context(run_context.id) target_execution = None # 定位到目标组件的Execution for exec in executions: exec_contexts = store.get_contexts_by_execution(exec.id) for ctx in exec_contexts: if ctx.type_name == 'node' and ctx.name.endswith(f'.{component_name}'): target_execution = exec break if target_execution: break if not target_execution: raise ValueError(f"Run {run_id}中未找到组件: {component_name}") # 获取该Execution产出的所有Artifact output_links = store.get_artifact_connections_by_execution( target_execution.id, mlmd.ExecutionEdge.OUTPUT ) artifact_ids = [link.artifact_id for link in output_links] artifacts = [store.get_artifact(aid) for aid in artifact_ids] # 筛选目标类型的Artifact artifact_type = store.get_artifact_type(type_name) if not artifact_type: raise ValueError(f"未找到Artifact类型: {type_name}") target_artifacts = [a for a in artifacts if a.type_id == artifact_type.id] # 转换为TFX Artifact对象返回 tfx_artifacts = [] for art in target_artifacts: tfx_art = Artifact(artifact_type) tfx_art.set_mlmd_artifact(art) tfx_artifacts.append(tfx_art) return tfx_artifacts
3. 使用示例
# 假设已获取本次运行的run_id sqlite_path = './simple_pipeline_data/metadata.db' connection_config = tfx.orchestration.metadata.sqlite_metadata_connection_config(sqlite_path) # 获取本次Run中SchemaGen生成的Schema Artifact schema_artifacts = get_artifacts_by_run_id( connection_config, run_id=run_id, component_name='SchemaGen', type_name='Schema' )
方案说明
- 该方案利用MLMD的Run Context机制,每个Pipeline运行的Run ID全局唯一,完全不受Pipeline名称或并发执行影响;
- 通过Run ID关联到本次运行的所有组件Execution,再精准定位目标组件产出的Artifact,确保获取的是本次运行的产物;
- MLMD原生支持此功能,不属于缺失特性。
内容的提问来源于stack exchange,提问作者Mehran
相关产品推荐
相关产品推荐

