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

Azure ML Pipeline执行Promptflow后如何正确聚合结果并在关联实验运行中记录MLflow指标

Azure ML Pipeline执行Promptflow后如何正确聚合结果并在关联实验运行中记录MLflow指标

我帮你分析下问题所在,再给出针对性的解决方案。你现在遇到的嵌套实验和指标无法正确关联的问题,核心是手动启动的MLflow父Run和Azure ML Pipeline自动创建的Run形成了嵌套关系,导致指标被记录到了手动启动的父Run里,而不是Pipeline对应的Run中。

问题根源拆解

你当前的代码里做了两件导致嵌套的操作:

  • 手动调用mlflow.start_run()创建了一个独立的MLflow父Run
  • 提交Pipeline时通过tags={"mlflow.parentRunId": parent_run.info.run_id},把Pipeline的Run设置为这个父Run的子Run
    最后在主脚本里执行mlflow.log_metric时,上下文还是那个手动启动的父Run,所以指标就跑到了嵌套的父Run中,而不是你期望的Pipeline自身的Run里。

方案一:在Pipeline中添加聚合步骤(推荐,符合MLOps规范)

这种方式把聚合逻辑做成Pipeline的一个专门步骤,指标会自动关联到Pipeline的Run上下文,不需要手动管理MLflow会话。

步骤1:创建聚合组件

新建aggregate_results文件夹,包含以下两个文件:

aggregate.py(聚合逻辑+MLflow指标记录)

import mlflow
import pandas as pd
import argparse

def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--input_data", type=str, help="Path to Promptflow output data")
    args = parser.parse_args()

    # 读取Promptflow输出的结果(根据你的flow输出格式调整,这里假设是CSV)
    df = pd.read_csv(args.input_data)
    
    # 执行你的聚合逻辑,这里用平均分作为示例
    avg_score = df["score"].mean()
    print(f"计算得到平均分数: {avg_score}")

    # 自动关联到Pipeline的Run上下文,直接记录指标
    mlflow.log_metric("Avg", avg_score)

if __name__ == "__main__":
    main()

component.yaml(定义Azure ML可复用组件)

$schema: https://azuremlschemas.azureedge.net/latest/commandComponent.schema.json
name: aggregate_promptflow_results
display_name: Aggregate Promptflow Results
version: 1
type: command
inputs:
  input_data:
    type: uri_file
outputs:
  aggregated_metrics:
    type: uri_file
code: ./
command: >-
  python aggregate.py
  --input_data ${{inputs.input_data}}
environment: azureml://registries/azureml/environments/sklearn-1.0/versions/3

步骤2:修改主Pipeline脚本

移除手动启动的MLflow父Run,添加聚合步骤到Pipeline中:

import os
import uuid
import mlflow
import azureml.mlflow
from azure.identity import DefaultAzureCredential
from azure.ai.ml import MLClient, load_component, Input, Output
from azure.ai.ml.constants import AssetTypes
from azure.ai.ml.dsl import pipeline

# ------------------------------------------------------------------------------
# 1. Workspace Configuration
# ------------------------------------------------------------------------------
subscription_id = "foo"
resource_group = "bar"
workspace_name = "baz"

os.environ['subscription_id'] = subscription_id
os.environ['resource_group'] = resource_group
os.environ['workspace_name'] = workspace_name

cred = DefaultAzureCredential()
ml_client = MLClient(
    credential = cred,
    subscription_id = subscription_id,
    resource_group_name = resource_group,
    workspace_name = workspace_name,
)

experiment_name = "test_custom_connection_promptflow_pipeline"

# ------------------------------------------------------------------------------
# 2. Load Components (Promptflow + Aggregate)
# ------------------------------------------------------------------------------
# 加载Promptflow组件
flow_component = load_component(source="flow.dag.yaml")

# 加载聚合组件
aggregate_component = load_component(source="aggregate_results/component.yaml")

# 注册Promptflow组件(版本自动递增)
existing_versions = ml_client.components.list(name="test_custom_connection")
existing_versions = [int(c.version) for c in existing_versions if c.version.isdigit()]
next_version = str(max(existing_versions) + 1) if existing_versions else "1"
ml_client.components.create_or_update(flow_component, version=str(next_version))

# ------------------------------------------------------------------------------
# 3. Build DSL Pipeline
# ------------------------------------------------------------------------------
local_csv_path = "sample.csv"
eval_data = Input(
    type=AssetTypes.URI_FILE,
    path=local_csv_path,
    mode="ro_mount",
)

@pipeline()
def eval_pipeline( ):
    # Promptflow执行步骤
    flow_node = flow_component(
        data=eval_data,
        topic="${data.topic}",
    )
    flow_node.compute = "LLM-Prompt-Flow"
    flow_node.resources = {"instance_count": 1}
    flow_node.mini_batch_size = 5
    flow_node.max_concurrency_per_instance = 1
    flow_node.error_threshold = -1
    flow_node.mini_batch_error_threshold = -1
    flow_node.logging_level = "DEBUG"

    # 聚合步骤:依赖Promptflow的输出(需和你的flow.dag.yaml中输出名称对应)
    aggregate_node = aggregate_component(
        input_data=flow_node.outputs.output
    )
    aggregate_node.compute = "LLM-Prompt-Flow"  # 可替换为轻量CPU集群节省成本

    return {
        "aggregated_metrics": aggregate_node.outputs.aggregated_metrics
    }

# 创建并提交Pipeline
pipeline_job = eval_pipeline()
pipeline_job.settings.default_compute = "LLM-Prompt-Flow"
pipeline_job.name = f"eval-{uuid.uuid4().hex[:8]}"

# 提交Pipeline,Azure ML会自动在实验下创建对应的Run
submitted = ml_client.jobs.create_or_update(
    pipeline_job,
    experiment_name=experiment_name
)
print(f"▶️ Submitted pipeline job: {submitted.name}")
ml_client.jobs.stream(submitted.name)

方案二:主脚本中关联Pipeline Run后记录指标

如果不想添加额外的Pipeline步骤,也可以在Pipeline运行完成后,主动关联到它的MLflow Run再记录指标:

# 移除所有手动启动的mlflow.start_run()代码

# ... 其他Pipeline构建、提交代码不变 ...

# 提交并等待Pipeline运行完成
submitted = ml_client.jobs.create_or_update(
    pipeline_job,
    experiment_name=experiment_name
)
print(f"▶️ Submitted pipeline job: {submitted.name}")
ml_client.jobs.stream(submitted.name)

# 获取Pipeline对应的MLflow Run ID
job_run = ml_client.jobs.get(submitted.name)
run_id = job_run.tags.get("mlflow.runId")

# 关联到Pipeline的Run并记录指标
if run_id:
    tracking_uri = ml_client.workspaces.get(workspace_name).mlflow_tracking_uri
    mlflow.set_tracking_uri(tracking_uri)
    with mlflow.start_run(run_id=run_id):
        # 替换为你的实际聚合逻辑:读取Promptflow输出文件计算指标
        avg_value = 2  # 占位符,替换为真实聚合结果
        mlflow.log_metric("Avg", avg_value)
else:
    print("无法获取Pipeline对应的MLflow Run ID,指标未记录")

关键注意事项

  1. 除非明确需要嵌套Run结构,否则不要手动调用mlflow.start_run(),Azure ML会自动管理Pipeline的Run上下文
  2. 方案一中的聚合步骤可以使用轻量CPU集群,不需要和Promptflow共用GPU集群,大幅节省成本
  3. 确保Promptflow组件的输出名称和聚合步骤的输入名称完全匹配(查看flow.dag.yaml中的outputs定义)

内容来源于stack exchange

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.08 11:15:35