Skip to content

Commit d7b8b37

Browse files
committed
Now working on GPUs
1 parent 6103dff commit d7b8b37

4 files changed

Lines changed: 298 additions & 9 deletions

File tree

‎main.py‎

Lines changed: 80 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,80 @@
1+
"""
2+
Distributed Tensorflow example
3+
The original code was in @ischlag, but the distributed architecture is quite
4+
different.
5+
The code runs on TF 1.1.
6+
Trains a simple sigmoid neural network on mnist for 20 epochs on three machines using one parameter server.
7+
The code requires 'tmux'.
8+
The code runs on the local server only.
9+
10+
Run like this:
11+
$ bash run.sh
12+
13+
Then, by using ctrl+b+(window number, e.g., 0, 1, 2),
14+
you can change the terminal.
15+
16+
"""
17+
from __future__ import print_function
18+
import tensorflow as tf
19+
import numpy as np
20+
import os
21+
import time
22+
import signal, sys
23+
from worker import Worker
24+
from utils import *
25+
26+
# Define flags.
27+
flags = tf.app.flags
28+
29+
flags.DEFINE_string('job_name', 'ps', "Either 'ps' or 'worker'")
30+
flags.DEFINE_integer('task_index', 0, "Index of task within the job")
31+
flags.DEFINE_integer('batch_size', 100, "Batch size")
32+
flags.DEFINE_float('learning_rate', 0.0005, "Learning rate")
33+
flags.DEFINE_integer('training_epochs', 20, "Training epochs")
34+
flags.DEFINE_string('logdir', './tmp/mnist/1', "Log directory")
35+
flags.DEFINE_integer('num_workers', 2, "Number of workers")
36+
flags.DEFINE_integer('num_gpus', 1,
37+
"Number of gpus, less than or equal to num_workers")
38+
39+
FLAGS = flags.FLAGS
40+
41+
def main():
42+
# Load MNIST dataset.
43+
from tensorflow.examples.tutorials.mnist import input_data
44+
mnist = input_data.read_data_sets('MNIST_data', one_hot=True)
45+
46+
# Cluster specification
47+
spec = cluster_spec(FLAGS.num_workers, 1)
48+
cluster = tf.train.ClusterSpec(spec)
49+
50+
# Signal
51+
def shutdown(signal, frame):
52+
sys.exit(128+signal)
53+
signal.signal(signal.SIGHUP, shutdown)
54+
signal.signal(signal.SIGINT, shutdown)
55+
signal.signal(signal.SIGTERM, shutdown)
56+
57+
# Set GPU memory fraction.
58+
process_per_memory =\
59+
np.ceil(float(FLAGS.num_workers)/float(FLAGS.num_gpus))
60+
fraction = 0.99 / process_per_memory
61+
gpu_options = tf.GPUOptions(per_process_gpu_memory_fraction=fraction)
62+
63+
if FLAGS.job_name == 'ps':
64+
config = tf.ConfigProto(device_filters=['/job:ps'])
65+
server = tf.train.Server(cluster, job_name='ps',
66+
task_index=FLAGS.task_index, config=config)
67+
while True:
68+
time.sleep(1000)
69+
70+
elif FLAGS.job_name == 'worker':
71+
config = tf.ConfigProto(gpu_options=gpu_options,
72+
intra_op_parallelism_threads=1,
73+
inter_op_parallelism_threads=2)
74+
server = tf.train.Server(cluster, job_name='worker',
75+
task_index=FLAGS.task_index, config=config)
76+
worker = Worker(FLAGS.job_name, FLAGS.task_index, server)
77+
worker.learn(mnist)
78+
79+
if __name__ == '__main__':
80+
main()

‎run.sh‎

Lines changed: 22 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -1,13 +1,26 @@
11
#!/bin/bash
2-
tmux kill-session -t examplesess
3-
tmux new-session -s examplesess -n ps-0 -d bash
4-
tmux new-window -t examplesess -n ps-1 -d bash
5-
tmux new-window -t examplesess -n worker-0 -d bash
6-
tmux new-window -t examplesess -n worker-1 -d bash
2+
name=sess
3+
num_workers=2
4+
num_gpus=1
5+
GPU_ID=(0)
6+
7+
tmux kill-session -t $name
8+
9+
tmux new-session -s $name -n ps -d bash
10+
for (( i=0; i<$num_workers; i++ ))
11+
do
12+
tmux new-window -t $name -n worker$i -d bash
13+
done
14+
715
sleep 1
8-
tmux send-keys -t examplesess:ps-0 'CUDA_VISIBLE_DEVICES= python example.py --job_name ps --task_index 0' Enter
9-
tmux send-keys -t examplesess:ps-1 'CUDA_VISIBLE_DEVICES= python example.py --job_name ps --task_index 1' Enter
10-
tmux send-keys -t examplesess:worker-0 'CUDA_VISIBLE_DEVICES=0 python example.py --job_name worker --task_index 0' Enter
11-
tmux send-keys -t examplesess:worker-1 'CUDA_VISIBLE_DEVICES=0 python example.py --job_name worker --task_index 1' Enter
16+
17+
tmux send-keys -t $name:ps "CUDA_VISIBLE_DEVICES= python main.py --num_workers $num_workers --num_gpus $num_gpus --job_name ps --task_index=0" Enter
18+
for (( i=0; i<$num_workers; i++ ))
19+
do
20+
ID=$((i % $num_gpus))
21+
tmux send-keys -t $name:worker$i "CUDA_VISIBLE_DEVICES=${GPU_ID[$ID]} python main.py --num_workers $num_workers --num_gpus $num_gpus --job_name worker --task_index $i" Enter
22+
done
23+
1224
sleep 1
25+
1326
tmux a

‎utils.py‎

Lines changed: 32 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,32 @@
1+
import tensorflow as tf
2+
3+
def get_vars(scope, trainable=True):
4+
if trainable:
5+
keys = tf.GraphKeys.TRAINABLE_VARIABLES
6+
else:
7+
keys = tf.GraphKeys.GLOBAL_VARIABLES
8+
return tf.get_collection(keys, scope)
9+
10+
def cluster_spec(num_workers, num_ps):
11+
cluster = {}
12+
port = 12222
13+
14+
all_ps = []
15+
host = '127.0.0.1'
16+
for _ in range(num_ps):
17+
all_ps.append('{}:{}'.format(host, port))
18+
port += 1
19+
cluster['ps'] = all_ps
20+
21+
all_workers = []
22+
for _ in range(num_workers):
23+
all_workers.append('{}:{}'.format(host, port))
24+
port += 1
25+
cluster['worker'] = all_workers
26+
return cluster
27+
28+
class FastSaver(tf.train.Saver):
29+
def save(self, sess, save_path, global_step=None, latest_filename=None,
30+
meta_graph_suffix='meta', write_meta_graph=True):
31+
super(FastSaver, self).save(sess, save_path, global_step,
32+
latest_filename, meta_graph_suffix, False)

‎worker.py‎

Lines changed: 164 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,164 @@
1+
import tensorflow as tf
2+
import time
3+
from utils import *
4+
5+
FLAGS = tf.app.flags.FLAGS
6+
7+
class Worker(object):
8+
def __init__(self, job_name, task_index, server):
9+
self.job_name = job_name
10+
self.task_index = task_index
11+
self.server = server
12+
13+
# For shared parameters, including global step.
14+
global_device = '/job:{}/task:{}/cpu:0'.format(job_name, task_index)
15+
16+
# For local computations,
17+
"""
18+
The gradient is computed at each "local_device".
19+
Since CUDA_VISIBLE_DEVICES for each worker process allocates single
20+
gpu, '/gpu:0' is used.
21+
"""
22+
local_device = '/job:{}/task:{}/gpu:0'.format(job_name, task_index)
23+
24+
with tf.device(tf.train.replica_device_setter(1,
25+
worker_device=global_device)):
26+
27+
with tf.variable_scope('global'):
28+
self.build_net()
29+
self.global_step = tf.get_variable('global_step', [], tf.int32,
30+
initializer=tf.constant_initializer(0, dtype=tf.int32),
31+
trainable=False)
32+
self.counter_op = self.global_step.assign_add(1)
33+
34+
with tf.device(local_device):
35+
with tf.variable_scope('local'):
36+
self.build_net()
37+
self.build_loss()
38+
self.build_sync_op()
39+
self.build_train_op()
40+
self.build_summary_op()
41+
42+
self.build_init_op()
43+
self.build_saver()
44+
45+
46+
def build_net(self):
47+
self.x = tf.placeholder(tf.float32, [None, 784])
48+
49+
def _net(inputs):
50+
net = tf.layers.dense(inputs, 100, activation=tf.nn.sigmoid,
51+
kernel_initializer=tf.random_normal_initializer())
52+
logits = tf.layers.dense(net, 10,
53+
kernel_initializer=tf.random_normal_initializer())
54+
net = tf.nn.softmax(logits)
55+
return net, logits
56+
57+
self.net, self.logits = _net(self.x)
58+
59+
def build_loss(self):
60+
self.y = tf.placeholder(tf.float32, [None, 10])
61+
62+
def _loss(labels, logits):
63+
cross_entropy = tf.nn.softmax_cross_entropy_with_logits(
64+
labels=labels, logits=logits)
65+
66+
return tf.reduce_mean(cross_entropy)
67+
68+
self.loss = _loss(self.y, self.logits)
69+
70+
def build_train_op(self):
71+
optimizer = tf.train.GradientDescentOptimizer(FLAGS.learning_rate)
72+
gvs = optimizer.compute_gradients(self.loss,
73+
var_list=get_vars('local'))
74+
75+
global_gvs = []
76+
for v, gv in zip(get_vars('global'), gvs):
77+
global_gvs.append((gv[0], v))
78+
79+
self.train_op = optimizer.apply_gradients(global_gvs)
80+
81+
def build_sync_op(self):
82+
local_vars = get_vars('local')
83+
global_vars = get_vars('global')
84+
self.sync_op = tf.group(*[v1.assign(v2)\
85+
for v1, v2 in zip(local_vars, global_vars)])
86+
87+
def build_summary_op(self):
88+
with tf.name_scope('accuracy'):
89+
correct_prediction = tf.equal(tf.argmax(self.net, 1), tf.argmax(self.y, 1))
90+
accuracy = self.accuracy = tf.reduce_mean(tf.cast(correct_prediction, tf.float32))
91+
92+
tf.summary.scalar('loss', self.loss)
93+
tf.summary.scalar('accuracy', accuracy)
94+
95+
self.summary_op = tf.summary.merge_all()
96+
self.summary_writer = tf.summary.FileWriter(FLAGS.logdir + '_%d' % self.task_index)
97+
98+
def build_init_op(self):
99+
self.global_init_op = tf.variables_initializer(get_vars('global', False))
100+
self.local_init_op = tf.variables_initializer(get_vars('local', False))
101+
102+
def build_saver(self):
103+
self.saver = FastSaver(get_vars('global', False))
104+
105+
def learn(self, dataset):
106+
107+
sv = tf.train.Supervisor(is_chief=(self.task_index==0),
108+
logdir=FLAGS.logdir,
109+
saver=self.saver,
110+
summary_op=None,
111+
summary_writer=self.summary_writer,
112+
ready_op=tf.report_uninitialized_variables(
113+
get_vars('global', False)),
114+
global_step=self.global_step,
115+
save_model_secs=30,
116+
save_summaries_secs=30,
117+
init_op=self.global_init_op,
118+
local_init_op=self.local_init_op)
119+
120+
config = tf.ConfigProto(allow_soft_placement=True,
121+
log_device_placement=True)
122+
123+
with sv.managed_session(self.server.target, config=config) as sess, sess.as_default():
124+
125+
begin_time = time.time()
126+
frequency = 100
127+
# perform training cycles
128+
start_time = time.time()
129+
130+
epoch = 0
131+
while not sv.should_stop() and epoch < FLAGS.training_epochs:
132+
# number of batches in one epoch
133+
batch_count = int(dataset.train.num_examples/FLAGS.batch_size)
134+
count = 0
135+
for i in range(batch_count):
136+
sess.run(self.sync_op)
137+
batch_x, batch_y = dataset.train.next_batch(FLAGS.batch_size)
138+
139+
# perform the operations we defined earlier on batch
140+
_, cost, summary, step = sess.run(
141+
[self.train_op, self.loss, self.summary_op, self.global_step],
142+
feed_dict={self.x: batch_x, self.y: batch_y})
143+
self.summary_writer.add_summary(summary, step)
144+
145+
count += 1
146+
if count % frequency == 0 or i+1 == batch_count:
147+
elapsed_time = time.time() - start_time
148+
start_time = time.time()
149+
print("Step: %d," % (step+1),
150+
" Epoch: %2d," % (epoch+1),
151+
" Batch: %3d of %3d," % (i+1, batch_count),
152+
" Cost: %.4f," % cost,
153+
" AvgTime: %3.2fms" % float(elapsed_time*1000/frequency))
154+
count = 0
155+
sess.run(self.counter_op)
156+
157+
epoch += 1
158+
159+
print("Test-Accuracy: %2.2f" % sess.run(self.accuracy,
160+
feed_dict={self.x: dataset.test.images, self.y: dataset.test.labels}))
161+
print("Total Time: %3.2fs" % float(time.time() - begin_time))
162+
print("Final Cost: %.4f" % cost)
163+
164+
print("done")

0 commit comments

Comments
 (0)