求基于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.
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.
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()
Let’s break down the critical changes you need to make compared to sync examples like cifar10_estimator:
- Distributed Strategy: Use
ParameterServerStrategyinstead ofMirroredStrategy(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
ClusterSpecwith PS and worker nodes, and start a dedicated server for each task. - RunConfig: Configure
train_distribute,master,task_type, andtask_idto inform the Estimator it’s operating in a distributed async environment.
- 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
- PS Node:
- You’ll see logs from each worker showing they’re training independently, updating the global model stored on the PS node.
To use this with your existing model (like CIFAR-10):
- Replace the
model_fnwith your existing model definition (keep the training op using a standard optimizer). - Swap the
input_fnto load your actual dataset (ensure your pipeline is scalable for distributed training). - Update the
cluster_specwith your real server IPs and ports.
内容的提问来源于stack exchange,提问作者Roy J.Swagger

