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

如何在TFX Kubeflow管道中同时运行多个Trainer组件?

Hey there! Great job getting your TFX pipeline up and running on Kubeflow on GCP—reusing components like you’re aiming to do is a smart move to keep your pipeline efficient and maintainable. Let’s walk through how you can reuse your CSVExampleGen, Transform, and SchemaGen while spinning up multiple Trainer components (each with their own outputs and pushers).

Core Idea

The key here is that TFX components are designed to pass their outputs downstream, so you can point multiple Trainer instances to the same outputs from your shared preprocessing components. With enable_cache=True, TFX will only run those shared components once, even if multiple trainers depend on them.

Modified Pipeline Code

Let’s adjust your create_pipeline function to support multiple training tasks. We’ll define a list of training configurations (each with its own run_fn, evaluation threshold, and serving directory) and loop through them to create independent Trainer, Evaluator, and Pusher components for each task.

import tensorflow_model_analysis as tfma
from tfx.components import (
    CsvExampleGen, SchemaGen, Transform, Trainer, Evaluator, Pusher
)
from tfx.proto import example_gen_pb2, trainer_pb2, pusher_pb2
from tfx.types import executor_spec
from tfx.components.trainer.executor import GenericExecutor
from typing import Text, Optional, List, Dict

def create_pipeline(
    pipeline_name: Text, 
    pipeline_root: Text, 
    data_path: Text, 
    shared_preprocessing_fn: Text,
    train_tasks: List[Dict],  # Each task holds config for a unique model
    train_args: trainer_pb2.TrainArgs, 
    eval_args: trainer_pb2.EvalArgs, 
    metadata_connection_config: Optional[metadata_store_pb2.ConnectionConfig] = None 
) -> pipeline.Pipeline:
    # ----------------------
    # Shared Base Components
    # ----------------------
    # Reused CSVExampleGen
    example_gen = CsvExampleGen(input_base=data_path)
    components = [example_gen]

    # Reused SchemaGen
    schema_gen = SchemaGen(statistics=example_gen.outputs['statistics'])
    components.append(schema_gen)

    # Reused Transform
    transform = Transform(
        examples=example_gen.outputs['examples'],
        schema=schema_gen.outputs['schema'],
        preprocessing_fn=shared_preprocessing_fn
    )
    components.append(transform)

    # ----------------------
    # Per-Task Components
    # ----------------------
    for task_idx, task_config in enumerate(train_tasks):
        # Unique identifier for this task's components (for clarity in UI/metadata)
        task_suffix = f"_model_{task_idx}"

        # 1. Create Trainer for this task
        trainer = Trainer(
            run_fn=task_config['run_fn'],
            transformed_examples=transform.outputs['transformed_examples'],
            schema=schema_gen.outputs['schema'],
            transform_graph=transform.outputs['transform_graph'],
            train_args=train_args,
            eval_args=eval_args,
            custom_executor_spec=executor_spec.ExecutorClassSpec(GenericExecutor),
            instance_name=f"trainer{task_suffix}"  # Unique name for tracking
        )
        components.append(trainer)

        # 2. Create Evaluator for this task's model
        eval_config = tfma.EvalConfig(
            model_specs=[tfma.ModelSpec(label_key="your_label_key")],
            metrics_specs=[tfma.MetricsSpec(
                metrics=[tfma.MetricConfig(class_name="BinaryAccuracy")]
            )],
            slicing_specs=[tfma.SlicingSpec()]
        )

        evaluator = Evaluator(
            examples=example_gen.outputs['examples'],
            model=trainer.outputs['model'],
            schema=schema_gen.outputs['schema'],
            eval_config=eval_config,
            threshold=example_gen_pb2.Threshold(
                value=task_config['eval_accuracy_threshold']
            ),
            instance_name=f"evaluator{task_suffix}"
        )
        components.append(evaluator)

        # 3. Create Pusher to deploy the blessed model
        pusher = Pusher(
            model=trainer.outputs['model'],
            model_blessing=evaluator.outputs['blessing'],
            push_destination=pusher_pb2.PushDestination(
                filesystem=pusher_pb2.PushDestination.Filesystem(
                    base_directory=task_config['serving_model_dir']
                )
            ),
            instance_name=f"pusher{task_suffix}"
        )
        components.append(pusher)

    # ----------------------
    # Assemble Pipeline
    # ----------------------
    return pipeline.Pipeline(
        pipeline_name=pipeline_name,
        pipeline_root=pipeline_root,
        components=components,
        enable_cache=True,  # Critical for reusing shared components
        metadata_connection_config=metadata_connection_config,
        beam_pipeline_args=beam_pipeline_args,
    )

How to Use This

When calling create_pipeline, pass a list of train_tasks where each entry defines a unique model:

train_tasks = [
    {
        "run_fn": "models.model_v1.run_fn",
        "eval_accuracy_threshold": 0.92,
        "serving_model_dir": "gs://your-bucket/models/v1"
    },
    {
        "run_fn": "models.model_v2.run_fn",
        "eval_accuracy_threshold": 0.94,
        "serving_model_dir": "gs://your-bucket/models/v2"
    }
]

# Initialize and run the pipeline
pipeline = create_pipeline(
    pipeline_name="multi_model_pipeline",
    pipeline_root="gs://your-bucket/pipeline_root",
    data_path="gs://your-bucket/data",
    shared_preprocessing_fn="preprocessing.preprocess_fn",
    train_tasks=train_tasks,
    train_args=trainer_pb2.TrainArgs(num_steps=1000),
    eval_args=trainer_pb2.EvalArgs(num_steps=200)
)

Key Notes

  • Cache Efficiency: With enable_cache=True, TFX will skip re-running CsvExampleGen, SchemaGen, and Transform after their first successful run—this saves a ton of time and resources.
  • Component Uniqueness: The instance_name parameter ensures each trainer/evaluator/pusher has a distinct identity in Kubeflow UI and TFX metadata, making it easier to debug and track each model’s progress.
  • Hyperparameter Tuning: If you want to test different hyperparameters for the same model architecture, add a hyperparameters key to each task_config and pass it to the Trainer component via the hyperparameters argument.
  • Resource Allocation: On Kubeflow, multiple trainers will run in parallel if your cluster has enough resources. Adjust your cluster’s node pool size if you need to handle more concurrent training jobs.

This setup lets you reuse all your preprocessing logic while experimenting with different model architectures, hyperparameters, or evaluation criteria—perfect for iterative model development!

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.07 18:17:32