TensorFlow分布式训练中模型保存与恢复问题求助
解决TensorFlow分布式训练中模型保存的问题
我之前在做TensorFlow 1.x分布式训练的时候也碰到过一模一样的问题!官方文档的示例确实没把模型保存这块讲透,尤其是在多worker+ps的架构下,很容易踩坑。结合你用的gcr.io/tensorflow/tensorflow:1.5.0-gpu-py3镜像,我给你梳理下可行的解决方案:
核心问题:避免多进程同时写模型
在分布式架构里,ps节点只负责参数存储,所有worker都在并行计算,如果每个worker都去保存模型,会导致文件写入冲突,而且完全没必要——我们只需要让**其中一个worker(通常选worker 0)**来完成模型保存操作。
具体代码调整
1. PS进程代码(保持参数服务逻辑即可)
import tensorflow as tf # 定义集群配置,替换成你的实际主机/端口 cluster = tf.train.ClusterSpec({ "ps": ["ps_container:2222"], "worker": ["worker0_container:2223", "worker1_container:2224"] }) # 创建PS服务器并阻塞等待 server = tf.train.Server(cluster, job_name="ps", task_index=0) server.join()
2. Worker进程代码(添加模型保存逻辑,仅worker 0执行)
基于MNIST示例调整,重点是通过task_index判断当前worker身份,只有0号worker负责模型保存:
import tensorflow as tf import os from tensorflow.examples.tutorials.mnist import input_data # 加载MNIST数据集 mnist = input_data.read_data_sets("./mnist_data", one_hot=True) # 集群配置,和PS端保持一致 cluster = tf.train.ClusterSpec({ "ps": ["ps_container:2222"], "worker": ["worker0_container:2223", "worker1_container:2224"] }) # 获取当前worker的任务索引,通过Docker环境变量传递 task_index = int(os.environ.get("TASK_INDEX", 0)) server = tf.train.Server(cluster, job_name="worker", task_index=task_index) # 构建分布式模型 with tf.device(tf.train.replica_device_setter( worker_device="/job:worker/task:%d" % task_index, cluster=cluster)): # MNIST基础模型定义 x = tf.placeholder(tf.float32, [None, 784]) W = tf.Variable(tf.zeros([784, 10])) b = tf.Variable(tf.zeros([10])) y = tf.matmul(x, W) + b y_ = tf.placeholder(tf.float32, [None, 10]) cross_entropy = tf.reduce_mean( tf.nn.softmax_cross_entropy_with_logits(labels=y_, logits=y)) # 全局步数,用于同步训练进度 global_step = tf.train.get_or_create_global_step() train_step = tf.train.GradientDescentOptimizer(0.5).minimize( cross_entropy, global_step=global_step) # 定义模型保存器 saver = tf.train.Saver() save_dir = "./saved_model" # 确保保存目录存在 tf.gfile.MakeDirs(save_dir) # 启动分布式训练会话 with tf.train.MonitoredTrainingSession( master=server.target, is_chief=(task_index == 0), # 指定worker0为chief,负责初始化和保存 checkpoint_dir=save_dir) as mon_sess: while not mon_sess.should_stop(): batch_xs, batch_ys = mnist.train.next_batch(100) mon_sess.run(train_step, feed_dict={x: batch_xs, y_: batch_ys}) # 训练结束后,由chief worker额外保存一次完整模型(可选) if task_index == 0: saver.save(mon_sess.raw_session(), save_dir + "/final_model.ckpt")
Docker运行注意事项
- 确保所有ps和worker容器在同一个Docker网络中,这样才能互相通信:
# 创建专属网络 docker network create tf-distributed-net - 启动容器时给worker传递
TASK_INDEX环境变量,区分不同worker节点:# 启动PS容器 docker run --name ps --network tf-distributed-net gcr.io/tensorflow/tensorflow:1.5.0-gpu-py3 python ps.py # 启动worker0(负责保存模型) docker run --name worker0 --network tf-distributed-net -e TASK_INDEX=0 gcr.io/tensorflow/tensorflow:1.5.0-gpu-py3 python worker.py # 启动worker1 docker run --name worker1 --network tf-distributed-net -e TASK_INDEX=1 gcr.io/tensorflow/tensorflow:1.5.0-gpu-py3 python worker.py
关键要点
- 使用
tf.train.MonitoredTrainingSession会自动帮你处理chief worker的初始化、checkpoint自动保存,比手动调用saver.save更稳妥 - 必须严格保证只有一个进程写入模型文件,否则会出现文件损坏、内容覆盖等问题
- 保存后的模型可以用
tf.train.Saver.restore在单机或分布式环境中加载使用
内容的提问来源于stack exchange,提问作者DWendt
相关产品推荐
相关产品推荐

