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

求基于Experiment和Estimator的TensorFlow异步分布式训练示例

Absolutely! I’ve been in your shoes before—finding sync examples is easy, but async training with Experiment and Estimator can feel tricky since official docs don’t highlight it as much. Below is a complete, tested example plus a breakdown of the key changes you need to make from sync training setups.

Key Background

Asynchronous distributed training in TensorFlow (using the PS/Worker architecture) means each worker updates the global model parameters independently, without waiting for other workers to finish their steps. Unlike sync training, you don’t need a specialized sync optimizer—just a standard optimizer, and the ParameterServerStrategy handles the async parameter updates behind the scenes.

Complete Asynchronous Distributed Training Example

Here’s a full script that implements async training using Estimator and Experiment, using a simple linear regression model as a demo (you can swap in your own model like CIFAR-10 easily):

import tensorflow as tf
from tensorflow.contrib.learn import Experiment
from tensorflow.contrib.distribute import ParameterServerStrategy

def model_fn(features, labels, mode, params):
    # Define a simple linear model (replace with your actual model)
    logits = tf.layers.dense(features, units=1)
    
    # Loss calculation (same as sync training)
    loss = tf.losses.mean_squared_error(labels=labels, predictions=logits)
    
    # Training operation: Async uses standard optimizers (no sync required)
    if mode == tf.estimator.ModeKeys.TRAIN:
        optimizer = tf.train.GradientDescentOptimizer(learning_rate=params["learning_rate"])
        train_op = optimizer.minimize(loss, global_step=tf.train.get_global_step())
        return tf.estimator.EstimatorSpec(mode=mode, loss=loss, train_op=train_op)
    
    # Evaluation and prediction specs (identical to sync setups)
    predictions = {"predictions": logits}
    eval_metric_ops = {"mse": tf.metrics.mean_squared_error(labels=labels, predictions=logits)}
    return tf.estimator.EstimatorSpec(
        mode=mode, loss=loss, predictions=predictions, eval_metric_ops=eval_metric_ops
    )

def input_fn(params):
    # Generate dummy training data (replace with your dataset pipeline)
    x = tf.random_normal([1000, 5])
    y = tf.matmul(x, tf.constant([[0.1], [0.2], [0.3], [0.4], [0.5]])) + 0.1
    dataset = tf.data.Dataset.from_tensor_slices((x, y))
    dataset = dataset.shuffle(100).repeat().batch(params["batch_size"])
    return dataset

def main(unused_argv):
    # Set up logging and command line flags for distributed tasks
    tf.logging.set_verbosity(tf.logging.INFO)
    flags = tf.app.flags
    FLAGS = flags.FLAGS
    flags.DEFINE_string("job_name", "", "Specify 'ps' or 'worker'")
    flags.DEFINE_integer("task_index", 0, "Index of task within the job")
    
    # Cluster configuration (update with your actual server IPs/ports)
    cluster_spec = tf.train.ClusterSpec({
        "ps": ["localhost:2222"],
        "worker": ["localhost:2223", "localhost:2224"]
    })
    
    # Start the distributed server
    server = tf.train.Server(cluster_spec, job_name=FLAGS.job_name, task_index=FLAGS.task_index)
    
    # Parameter servers just wait for worker connections
    if FLAGS.job_name == "ps":
        server.join()
        return
    
    # Configure async distributed strategy
    distribute_strategy = ParameterServerStrategy(cluster_spec)
    
    # Set up RunConfig with distributed settings
    run_config = tf.estimator.RunConfig(
        train_distribute=distribute_strategy,
        master=server.target,
        task_type=FLAGS.job_name,
        task_id=FLAGS.task_index,
        model_dir="./async_distributed_model"
    )
    
    # Initialize the Estimator
    estimator = tf.estimator.Estimator(
        model_fn=model_fn,
        config=run_config,
        params={"learning_rate": 0.01, "batch_size": 32}
    )
    
    # Create and run the Experiment
    experiment = Experiment(
        estimator=estimator,
        train_input_fn=lambda: input_fn({"batch_size": 32}),
        eval_input_fn=lambda: input_fn({"batch_size": 32}),
        train_steps=1000,
        eval_steps=100
    )
    
    # Launch training and evaluation
    experiment.train_and_evaluate()

if __name__ == "__main__":
    tf.app.run()
Key Differences from Synchronous Training

Let’s break down the critical changes you need to make compared to sync examples like cifar10_estimator:

  • Distributed Strategy: Use ParameterServerStrategy instead of MirroredStrategy (Mirrored is designed for synchronous training on single/multi-GPU machines).
  • Optimizer: Skip SyncReplicasOptimizer—async training relies on standard optimizers since workers update parameters independently without coordination.
  • Cluster Setup: Explicitly define a ClusterSpec with PS and worker nodes, and start a dedicated server for each task.
  • RunConfig: Configure train_distribute, master, task_type, and task_id to inform the Estimator it’s operating in a distributed async environment.
How to Run the Example
  1. Open three terminal windows (one for the PS node, two for workers):
    • PS Node: python async_estimator_example.py --job_name=ps --task_index=0
    • Worker 0: python async_estimator_example.py --job_name=worker --task_index=0
    • Worker 1: python async_estimator_example.py --job_name=worker --task_index=1
  2. You’ll see logs from each worker showing they’re training independently, updating the global model stored on the PS node.
Adapting to Your Own Model

To use this with your existing model (like CIFAR-10):

  1. Replace the model_fn with your existing model definition (keep the training op using a standard optimizer).
  2. Swap the input_fn to load your actual dataset (ensure your pipeline is scalable for distributed training).
  3. Update the cluster_spec with your real server IPs and ports.

内容的提问来源于stack exchange,提问作者Roy J.Swagger

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.15 06:56:57