##### Copyright 2019 The TensorFlow Authors.


In [1]:
#@title Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# https://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

# Distributed training with Keras

<table class="tfo-notebook-buttons" align="left">
  <td>
    <a target="_blank" href="https://www.tensorflow.org/tutorials/distribute/keras"><img src="https://www.tensorflow.org/images/tf_logo_32px.png" />View on TensorFlow.org</a>
  </td>
  <td>
    <a target="_blank" href="https://colab.research.google.com/github/tensorflow/docs/blob/master/site/en/tutorials/distribute/keras.ipynb"><img src="https://www.tensorflow.org/images/colab_logo_32px.png" />Run in Google Colab</a>
  </td>
  <td>
    <a target="_blank" href="https://github.com/tensorflow/docs/blob/master/site/en/tutorials/distribute/keras.ipynb"><img src="https://www.tensorflow.org/images/GitHub-Mark-32px.png" />View source on GitHub</a>
  </td>
  <td>
    <a href="https://storage.googleapis.com/tensorflow_docs/docs/site/en/tutorials/distribute/keras.ipynb"><img src="https://www.tensorflow.org/images/download_logo_32px.png" />Download notebook</a>
  </td>
</table>

## Overview

The `tf.distribute.Strategy` API provides an abstraction for distributing your training
across multiple processing units. The goal is to allow users to enable distributed training using existing models and training code, with minimal changes.

This tutorial uses the `tf.distribute.MirroredStrategy`, which
does in-graph replication with synchronous training on many GPUs on one machine.
Essentially, it copies all of the model's variables to each processor.
Then, it uses [all-reduce](http://mpitutorial.com/tutorials/mpi-reduce-and-allreduce/) to combine the gradients from all processors and applies the combined value to all copies of the model.

`MirroredStrategy` is one of several distribution strategy available in TensorFlow core. You can read about more strategies at [distribution strategy guide](../../guide/distributed_training.ipynb).


### Keras API

This example uses the `tf.keras` API to build the model and training loop. For custom training loops, see the [tf.distribute.Strategy with training loops](training_loops.ipynb) tutorial.

## Import dependencies

In [2]:
# Import TensorFlow and TensorFlow Datasets

import tensorflow_datasets as tfds
import tensorflow as tf
tfds.disable_progress_bar()

import os

In [3]:
print(tf.__version__)

2.2.0


## Download the dataset

Download the MNIST dataset and load it from [TensorFlow Datasets](https://www.tensorflow.org/datasets). This returns a dataset in `tf.data` format.

Setting `with_info` to `True` includes the metadata for the entire dataset, which is being saved here to `info`.
Among other things, this metadata object includes the number of train and test examples. 


In [4]:
datasets, info = tfds.load(name='mnist', with_info=True, as_supervised=True)

mnist_train, mnist_test = datasets['train'], datasets['test']

## Define distribution strategy

Create a `MirroredStrategy` object. This will handle distribution, and provides a context manager (`tf.distribute.MirroredStrategy.scope`) to build your model inside.

In [5]:
strategy = tf.distribute.MirroredStrategy()

INFO:tensorflow:Using MirroredStrategy with devices ('/job:localhost/replica:0/task:0/device:GPU:0',)


INFO:tensorflow:Using MirroredStrategy with devices ('/job:localhost/replica:0/task:0/device:GPU:0',)


In [6]:
print('Number of devices: {}'.format(strategy.num_replicas_in_sync))

Number of devices: 1


## Setup input pipeline

When training a model with multiple GPUs, you can use the extra computing power effectively by increasing the batch size. In general, use the largest batch size that fits the GPU memory, and tune the learning rate accordingly.

In [7]:
# You can also do info.splits.total_num_examples to get the total
# number of examples in the dataset.

num_train_examples = info.splits['train'].num_examples
num_test_examples = info.splits['test'].num_examples

BUFFER_SIZE = 10000

BATCH_SIZE_PER_REPLICA = 64
BATCH_SIZE = BATCH_SIZE_PER_REPLICA * strategy.num_replicas_in_sync

Pixel values, which are 0-255, [have to be normalized to the 0-1 range](https://en.wikipedia.org/wiki/Feature_scaling). Define this scale in a function.

In [8]:
def scale(image, label):
  image = tf.cast(image, tf.float32)
  image /= 255

  return image, label

Apply this function to the training and test data, shuffle the training data, and [batch it for training](https://www.tensorflow.org/api_docs/python/tf/data/Dataset#batch). Notice we are also keeping an in-memory cache of the training data to improve performance.


In [9]:
train_dataset = mnist_train.map(scale).cache().shuffle(BUFFER_SIZE).batch(BATCH_SIZE)
eval_dataset = mnist_test.map(scale).batch(BATCH_SIZE)

## Create the model

Create and compile the Keras model in the context of `strategy.scope`.

In [10]:
with strategy.scope():
  model = tf.keras.Sequential([
      tf.keras.layers.Conv2D(32, 3, activation='relu', input_shape=(28, 28, 1)),
      tf.keras.layers.MaxPooling2D(),
      tf.keras.layers.Flatten(),
      tf.keras.layers.Dense(64, activation='relu'),
      tf.keras.layers.Dense(10)
  ])

  model.compile(loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True),
                optimizer=tf.keras.optimizers.Adam(),
                metrics=['accuracy'])

## Define the callbacks


The callbacks used here are:

*   *TensorBoard*: This callback writes a log for TensorBoard which allows you to visualize the graphs.
*   *Model Checkpoint*: This callback saves the model after every epoch.
*   *Learning Rate Scheduler*: Using this callback, you can schedule the learning rate to change after every epoch/batch.

For illustrative purposes, add a print callback to display the *learning rate* in the notebook.

In [11]:
# Define the checkpoint directory to store the checkpoints

checkpoint_dir = './training_checkpoints'
# Name of the checkpoint files
checkpoint_prefix = os.path.join(checkpoint_dir, "ckpt_{epoch}")

In [12]:
# Function for decaying the learning rate.
# You can define any decay function you need.
def decay(epoch):
  if epoch < 3:
    return 1e-3
  elif epoch >= 3 and epoch < 7:
    return 1e-4
  else:
    return 1e-5

In [13]:
# Callback for printing the LR at the end of each epoch.
class PrintLR(tf.keras.callbacks.Callback):
  def on_epoch_end(self, epoch, logs=None):
    print('\nLearning rate for epoch {} is {}'.format(epoch + 1,
                                                      model.optimizer.lr.numpy()))

In [14]:
callbacks = [
    tf.keras.callbacks.TensorBoard(log_dir='./logs'),
    tf.keras.callbacks.ModelCheckpoint(filepath=checkpoint_prefix,
                                       save_weights_only=True),
    tf.keras.callbacks.LearningRateScheduler(decay),
    PrintLR()
]

## Train and evaluate

Now, train the model in the usual way, calling `fit` on the model and passing in the dataset created at the beginning of the tutorial. This step is the same whether you are distributing the training or not.


In [15]:
model.fit(train_dataset, epochs=12, callbacks=callbacks)

Epoch 1/12


INFO:tensorflow:Reduce to /job:localhost/replica:0/task:0/device:CPU:0 then broadcast to ('/job:localhost/replica:0/task:0/device:CPU:0',).


INFO:tensorflow:Reduce to /job:localhost/replica:0/task:0/device:CPU:0 then broadcast to ('/job:localhost/replica:0/task:0/device:CPU:0',).


INFO:tensorflow:Reduce to /job:localhost/replica:0/task:0/device:CPU:0 then broadcast to ('/job:localhost/replica:0/task:0/device:CPU:0',).


INFO:tensorflow:Reduce to /job:localhost/replica:0/task:0/device:CPU:0 then broadcast to ('/job:localhost/replica:0/task:0/device:CPU:0',).


INFO:tensorflow:Reduce to /job:localhost/replica:0/task:0/device:CPU:0 then broadcast to ('/job:localhost/replica:0/task:0/device:CPU:0',).


INFO:tensorflow:Reduce to /job:localhost/replica:0/task:0/device:CPU:0 then broadcast to ('/job:localhost/replica:0/task:0/device:CPU:0',).


INFO:tensorflow:Reduce to /job:localhost/replica:0/task:0/device:CPU:0 then broadcast to ('/job:localhost/replica:0/task:0/device:CPU:0',).


INFO:tensorflow:Reduce to /job:localhost/replica:0/task:0/device:CPU:0 then broadcast to ('/job:localhost/replica:0/task:0/device:CPU:0',).


INFO:tensorflow:Reduce to /job:localhost/replica:0/task:0/device:CPU:0 then broadcast to ('/job:localhost/replica:0/task:0/device:CPU:0',).


INFO:tensorflow:Reduce to /job:localhost/replica:0/task:0/device:CPU:0 then broadcast to ('/job:localhost/replica:0/task:0/device:CPU:0',).


INFO:tensorflow:Reduce to /job:localhost/replica:0/task:0/device:CPU:0 then broadcast to ('/job:localhost/replica:0/task:0/device:CPU:0',).


INFO:tensorflow:Reduce to /job:localhost/replica:0/task:0/device:CPU:0 then broadcast to ('/job:localhost/replica:0/task:0/device:CPU:0',).


INFO:tensorflow:Reduce to /job:localhost/replica:0/task:0/device:CPU:0 then broadcast to ('/job:localhost/replica:0/task:0/device:CPU:0',).


INFO:tensorflow:Reduce to /job:localhost/replica:0/task:0/device:CPU:0 then broadcast to ('/job:localhost/replica:0/task:0/device:CPU:0',).


INFO:tensorflow:Reduce to /job:localhost/replica:0/task:0/device:CPU:0 then broadcast to ('/job:localhost/replica:0/task:0/device:CPU:0',).


INFO:tensorflow:Reduce to /job:localhost/replica:0/task:0/device:CPU:0 then broadcast to ('/job:localhost/replica:0/task:0/device:CPU:0',).


  1/938 [..............................] - ETA: 0s - accuracy: 0.1250 - loss: 2.3156

  8/938 [..............................] - ETA: 6s - accuracy: 0.4258 - loss: 2.0437

 18/938 [..............................] - ETA: 5s - accuracy: 0.5764 - loss: 1.6732

 27/938 [..............................] - ETA: 5s - accuracy: 0.6424 - loss: 1.3997

 36/938 [>.............................] - ETA: 5s - accuracy: 0.6901 - loss: 1.1996

 45/938 [>.............................] - ETA: 5s - accuracy: 0.7222 - loss: 1.0519

 54/938 [>.............................] - ETA: 5s - accuracy: 0.7457 - loss: 0.9517

 63/938 [=>............................] - ETA: 5s - accuracy: 0.7679 - loss: 0.8630

 72/938 [=>............................] - ETA: 5s - accuracy: 0.7845 - loss: 0.7952

 81/938 [=>............................] - ETA: 5s - accuracy: 0.7965 - loss: 0.7431

 90/938 [=>............................] - ETA: 4s - accuracy: 0.8035 - loss: 0.7073

 99/938 [==>...........................] - ETA: 4s - accuracy: 0.8138 - loss: 0.6716

108/938 [==>...........................] - ETA: 4s - accuracy: 0.8215 - loss: 0.6433

117/938 [==>...........................] - ETA: 4s - accuracy: 0.8289 - loss: 0.6151

126/938 [===>..........................] - ETA: 4s - accuracy: 0.8346 - loss: 0.5909

136/938 [===>..........................] - ETA: 4s - accuracy: 0.8400 - loss: 0.5706

145/938 [===>..........................] - ETA: 4s - accuracy: 0.8443 - loss: 0.5552

154/938 [===>..........................] - ETA: 4s - accuracy: 0.8487 - loss: 0.5396

162/938 [====>.........................] - ETA: 4s - accuracy: 0.8531 - loss: 0.5253

171/938 [====>.........................] - ETA: 4s - accuracy: 0.8571 - loss: 0.5109

180/938 [====>.........................] - ETA: 4s - accuracy: 0.8605 - loss: 0.4971

190/938 [=====>........................] - ETA: 4s - accuracy: 0.8644 - loss: 0.4843

199/938 [=====>........................] - ETA: 4s - accuracy: 0.8687 - loss: 0.4710

208/938 [=====>........................] - ETA: 4s - accuracy: 0.8723 - loss: 0.4584

217/938 [=====>........................] - ETA: 4s - accuracy: 0.8754 - loss: 0.4462




















































































































































Learning rate for epoch 1 is 0.0010000000474974513


Epoch 2/12
  1/938 [..............................] - ETA: 0s - accuracy: 0.9531 - loss: 0.1005

 14/938 [..............................] - ETA: 3s - accuracy: 0.9710 - loss: 0.0966

 27/938 [..............................] - ETA: 3s - accuracy: 0.9745 - loss: 0.0876

 40/938 [>.............................] - ETA: 3s - accuracy: 0.9766 - loss: 0.0854

 53/938 [>.............................] - ETA: 3s - accuracy: 0.9785 - loss: 0.0829

 66/938 [=>............................] - ETA: 3s - accuracy: 0.9777 - loss: 0.0823

 80/938 [=>............................] - ETA: 3s - accuracy: 0.9783 - loss: 0.0816

 94/938 [==>...........................] - ETA: 3s - accuracy: 0.9791 - loss: 0.0779

108/938 [==>...........................] - ETA: 3s - accuracy: 0.9786 - loss: 0.0777

122/938 [==>...........................] - ETA: 3s - accuracy: 0.9786 - loss: 0.0779

135/938 [===>..........................] - ETA: 3s - accuracy: 0.9785 - loss: 0.0767

148/938 [===>..........................] - ETA: 3s - accuracy: 0.9786 - loss: 0.0759

161/938 [====>.........................] - ETA: 3s - accuracy: 0.9791 - loss: 0.0746

174/938 [====>.........................] - ETA: 2s - accuracy: 0.9793 - loss: 0.0734

188/938 [=====>........................] - ETA: 2s - accuracy: 0.9797 - loss: 0.0719

202/938 [=====>........................] - ETA: 2s - accuracy: 0.9794 - loss: 0.0722

215/938 [=====>........................] - ETA: 2s - accuracy: 0.9794 - loss: 0.0716














































































































Learning rate for epoch 2 is 0.0010000000474974513


Epoch 3/12
  1/938 [..............................] - ETA: 0s - accuracy: 1.0000 - loss: 0.0195

 16/938 [..............................] - ETA: 3s - accuracy: 0.9883 - loss: 0.0284

 31/938 [..............................] - ETA: 3s - accuracy: 0.9889 - loss: 0.0339

 46/938 [>.............................] - ETA: 3s - accuracy: 0.9881 - loss: 0.0422

 61/938 [>.............................] - ETA: 3s - accuracy: 0.9882 - loss: 0.0434

 76/938 [=>............................] - ETA: 2s - accuracy: 0.9879 - loss: 0.0455

 90/938 [=>............................] - ETA: 2s - accuracy: 0.9872 - loss: 0.0464

105/938 [==>...........................] - ETA: 2s - accuracy: 0.9859 - loss: 0.0492

120/938 [==>...........................] - ETA: 2s - accuracy: 0.9850 - loss: 0.0519

135/938 [===>..........................] - ETA: 2s - accuracy: 0.9837 - loss: 0.0557

150/938 [===>..........................] - ETA: 2s - accuracy: 0.9836 - loss: 0.0553

165/938 [====>.........................] - ETA: 2s - accuracy: 0.9837 - loss: 0.0547

180/938 [====>.........................] - ETA: 2s - accuracy: 0.9839 - loss: 0.0549

195/938 [=====>........................] - ETA: 2s - accuracy: 0.9841 - loss: 0.0547

209/938 [=====>........................] - ETA: 2s - accuracy: 0.9842 - loss: 0.0547




































































































Learning rate for epoch 3 is 0.0010000000474974513


Epoch 4/12
  1/938 [..............................] - ETA: 0s - accuracy: 0.9688 - loss: 0.0878

 15/938 [..............................] - ETA: 3s - accuracy: 0.9885 - loss: 0.0365

 30/938 [..............................] - ETA: 3s - accuracy: 0.9896 - loss: 0.0330

 45/938 [>.............................] - ETA: 3s - accuracy: 0.9903 - loss: 0.0331

 60/938 [>.............................] - ETA: 3s - accuracy: 0.9914 - loss: 0.0303

 75/938 [=>............................] - ETA: 2s - accuracy: 0.9910 - loss: 0.0302

 90/938 [=>............................] - ETA: 2s - accuracy: 0.9901 - loss: 0.0329

105/938 [==>...........................] - ETA: 2s - accuracy: 0.9894 - loss: 0.0359

119/938 [==>...........................] - ETA: 2s - accuracy: 0.9900 - loss: 0.0349

134/938 [===>..........................] - ETA: 2s - accuracy: 0.9902 - loss: 0.0341

149/938 [===>..........................] - ETA: 2s - accuracy: 0.9904 - loss: 0.0332

164/938 [====>.........................] - ETA: 2s - accuracy: 0.9904 - loss: 0.0324

179/938 [====>.........................] - ETA: 2s - accuracy: 0.9900 - loss: 0.0330

194/938 [=====>........................] - ETA: 2s - accuracy: 0.9903 - loss: 0.0323

209/938 [=====>........................] - ETA: 2s - accuracy: 0.9907 - loss: 0.0322


































































































Learning rate for epoch 4 is 9.999999747378752e-05


Epoch 5/12


  1/938 [..............................] - ETA: 0s - accuracy: 1.0000 - loss: 0.0046

 16/938 [..............................] - ETA: 3s - accuracy: 0.9932 - loss: 0.0202

 31/938 [..............................] - ETA: 3s - accuracy: 0.9950 - loss: 0.0159

 45/938 [>.............................] - ETA: 3s - accuracy: 0.9944 - loss: 0.0191

 59/938 [>.............................] - ETA: 3s - accuracy: 0.9934 - loss: 0.0241

 74/938 [=>............................] - ETA: 3s - accuracy: 0.9937 - loss: 0.0235

 88/938 [=>............................] - ETA: 2s - accuracy: 0.9931 - loss: 0.0253

103/938 [==>...........................] - ETA: 2s - accuracy: 0.9924 - loss: 0.0268

118/938 [==>...........................] - ETA: 2s - accuracy: 0.9926 - loss: 0.0263

133/938 [===>..........................] - ETA: 2s - accuracy: 0.9926 - loss: 0.0255

148/938 [===>..........................] - ETA: 2s - accuracy: 0.9929 - loss: 0.0244

163/938 [====>.........................] - ETA: 2s - accuracy: 0.9928 - loss: 0.0248

178/938 [====>.........................] - ETA: 2s - accuracy: 0.9930 - loss: 0.0243

193/938 [=====>........................] - ETA: 2s - accuracy: 0.9930 - loss: 0.0244

207/938 [=====>........................] - ETA: 2s - accuracy: 0.9934 - loss: 0.0236




































































































Learning rate for epoch 5 is 9.999999747378752e-05


Epoch 6/12
  1/938 [..............................] - ETA: 0s - accuracy: 0.9844 - loss: 0.0561

 16/938 [..............................] - ETA: 3s - accuracy: 0.9912 - loss: 0.0241

 30/938 [..............................] - ETA: 3s - accuracy: 0.9922 - loss: 0.0346

 45/938 [>.............................] - ETA: 3s - accuracy: 0.9944 - loss: 0.0275

 60/938 [>.............................] - ETA: 3s - accuracy: 0.9943 - loss: 0.0276

 75/938 [=>............................] - ETA: 3s - accuracy: 0.9944 - loss: 0.0254

 90/938 [=>............................] - ETA: 2s - accuracy: 0.9948 - loss: 0.0239

105/938 [==>...........................] - ETA: 2s - accuracy: 0.9949 - loss: 0.0227

120/938 [==>...........................] - ETA: 2s - accuracy: 0.9944 - loss: 0.0256

135/938 [===>..........................] - ETA: 2s - accuracy: 0.9939 - loss: 0.0258

150/938 [===>..........................] - ETA: 2s - accuracy: 0.9932 - loss: 0.0270

165/938 [====>.........................] - ETA: 2s - accuracy: 0.9929 - loss: 0.0272

180/938 [====>.........................] - ETA: 2s - accuracy: 0.9931 - loss: 0.0263

195/938 [=====>........................] - ETA: 2s - accuracy: 0.9930 - loss: 0.0263

209/938 [=====>........................] - ETA: 2s - accuracy: 0.9930 - loss: 0.0260




































































































Learning rate for epoch 6 is 9.999999747378752e-05


Epoch 7/12
  1/938 [..............................] - ETA: 0s - accuracy: 0.9844 - loss: 0.0219

 15/938 [..............................] - ETA: 3s - accuracy: 0.9875 - loss: 0.0281

 30/938 [..............................] - ETA: 3s - accuracy: 0.9917 - loss: 0.0242

 45/938 [>.............................] - ETA: 3s - accuracy: 0.9917 - loss: 0.0238

 60/938 [>.............................] - ETA: 3s - accuracy: 0.9922 - loss: 0.0240

 75/938 [=>............................] - ETA: 3s - accuracy: 0.9933 - loss: 0.0212

 89/938 [=>............................] - ETA: 2s - accuracy: 0.9933 - loss: 0.0209

104/938 [==>...........................] - ETA: 2s - accuracy: 0.9928 - loss: 0.0243

119/938 [==>...........................] - ETA: 2s - accuracy: 0.9925 - loss: 0.0243

134/938 [===>..........................] - ETA: 2s - accuracy: 0.9929 - loss: 0.0237

149/938 [===>..........................] - ETA: 2s - accuracy: 0.9929 - loss: 0.0246

163/938 [====>.........................] - ETA: 2s - accuracy: 0.9927 - loss: 0.0250

178/938 [====>.........................] - ETA: 2s - accuracy: 0.9930 - loss: 0.0242

193/938 [=====>........................] - ETA: 2s - accuracy: 0.9931 - loss: 0.0240

208/938 [=====>........................] - ETA: 2s - accuracy: 0.9931 - loss: 0.0239




































































































Learning rate for epoch 7 is 9.999999747378752e-05


Epoch 8/12
  1/938 [..............................] - ETA: 0s - accuracy: 0.9844 - loss: 0.0215

 16/938 [..............................] - ETA: 3s - accuracy: 0.9922 - loss: 0.0254

 31/938 [..............................] - ETA: 3s - accuracy: 0.9945 - loss: 0.0213

 45/938 [>.............................] - ETA: 3s - accuracy: 0.9944 - loss: 0.0241

 60/938 [>.............................] - ETA: 3s - accuracy: 0.9945 - loss: 0.0221

 75/938 [=>............................] - ETA: 3s - accuracy: 0.9942 - loss: 0.0228

 89/938 [=>............................] - ETA: 2s - accuracy: 0.9933 - loss: 0.0235

104/938 [==>...........................] - ETA: 2s - accuracy: 0.9931 - loss: 0.0224

119/938 [==>...........................] - ETA: 2s - accuracy: 0.9934 - loss: 0.0214

134/938 [===>..........................] - ETA: 2s - accuracy: 0.9938 - loss: 0.0215

149/938 [===>..........................] - ETA: 2s - accuracy: 0.9935 - loss: 0.0218

164/938 [====>.........................] - ETA: 2s - accuracy: 0.9938 - loss: 0.0210

179/938 [====>.........................] - ETA: 2s - accuracy: 0.9938 - loss: 0.0212

193/938 [=====>........................] - ETA: 2s - accuracy: 0.9939 - loss: 0.0212

208/938 [=====>........................] - ETA: 2s - accuracy: 0.9943 - loss: 0.0207




































































































Learning rate for epoch 8 is 9.999999747378752e-06


Epoch 9/12
  1/938 [..............................] - ETA: 0s - accuracy: 0.9844 - loss: 0.0355

 16/938 [..............................] - ETA: 3s - accuracy: 0.9961 - loss: 0.0202

 31/938 [..............................] - ETA: 3s - accuracy: 0.9940 - loss: 0.0208

 46/938 [>.............................] - ETA: 3s - accuracy: 0.9935 - loss: 0.0214

 61/938 [>.............................] - ETA: 3s - accuracy: 0.9928 - loss: 0.0256

 76/938 [=>............................] - ETA: 2s - accuracy: 0.9936 - loss: 0.0241

 91/938 [=>............................] - ETA: 2s - accuracy: 0.9936 - loss: 0.0238

106/938 [==>...........................] - ETA: 2s - accuracy: 0.9938 - loss: 0.0226

120/938 [==>...........................] - ETA: 2s - accuracy: 0.9943 - loss: 0.0213

135/938 [===>..........................] - ETA: 2s - accuracy: 0.9940 - loss: 0.0214

150/938 [===>..........................] - ETA: 2s - accuracy: 0.9936 - loss: 0.0213

165/938 [====>.........................] - ETA: 2s - accuracy: 0.9935 - loss: 0.0215

180/938 [====>.........................] - ETA: 2s - accuracy: 0.9938 - loss: 0.0211

195/938 [=====>........................] - ETA: 2s - accuracy: 0.9942 - loss: 0.0204

210/938 [=====>........................] - ETA: 2s - accuracy: 0.9943 - loss: 0.0206




































































































Learning rate for epoch 9 is 9.999999747378752e-06


Epoch 10/12
  1/938 [..............................] - ETA: 0s - accuracy: 1.0000 - loss: 0.0029

 15/938 [..............................] - ETA: 3s - accuracy: 0.9937 - loss: 0.0250

 29/938 [..............................] - ETA: 3s - accuracy: 0.9952 - loss: 0.0207

 44/938 [>.............................] - ETA: 3s - accuracy: 0.9961 - loss: 0.0171

 59/938 [>.............................] - ETA: 3s - accuracy: 0.9955 - loss: 0.0172

 74/938 [=>............................] - ETA: 3s - accuracy: 0.9951 - loss: 0.0174

 89/938 [=>............................] - ETA: 2s - accuracy: 0.9951 - loss: 0.0189

104/938 [==>...........................] - ETA: 2s - accuracy: 0.9941 - loss: 0.0224

119/938 [==>...........................] - ETA: 2s - accuracy: 0.9944 - loss: 0.0225

134/938 [===>..........................] - ETA: 2s - accuracy: 0.9944 - loss: 0.0223

149/938 [===>..........................] - ETA: 2s - accuracy: 0.9945 - loss: 0.0219

163/938 [====>.........................] - ETA: 2s - accuracy: 0.9946 - loss: 0.0213

178/938 [====>.........................] - ETA: 2s - accuracy: 0.9950 - loss: 0.0204

193/938 [=====>........................] - ETA: 2s - accuracy: 0.9952 - loss: 0.0201

208/938 [=====>........................] - ETA: 2s - accuracy: 0.9953 - loss: 0.0197




































































































Learning rate for epoch 10 is 9.999999747378752e-06


Epoch 11/12
  1/938 [..............................] - ETA: 0s - accuracy: 0.9688 - loss: 0.0275

 15/938 [..............................] - ETA: 3s - accuracy: 0.9948 - loss: 0.0148

 30/938 [..............................] - ETA: 3s - accuracy: 0.9948 - loss: 0.0167

 44/938 [>.............................] - ETA: 3s - accuracy: 0.9950 - loss: 0.0201

 59/938 [>.............................] - ETA: 3s - accuracy: 0.9958 - loss: 0.0179

 74/938 [=>............................] - ETA: 3s - accuracy: 0.9945 - loss: 0.0200

 89/938 [=>............................] - ETA: 2s - accuracy: 0.9939 - loss: 0.0228

103/938 [==>...........................] - ETA: 2s - accuracy: 0.9944 - loss: 0.0214

118/938 [==>...........................] - ETA: 2s - accuracy: 0.9946 - loss: 0.0216

133/938 [===>..........................] - ETA: 2s - accuracy: 0.9947 - loss: 0.0215

148/938 [===>..........................] - ETA: 2s - accuracy: 0.9945 - loss: 0.0217

163/938 [====>.........................] - ETA: 2s - accuracy: 0.9946 - loss: 0.0211

178/938 [====>.........................] - ETA: 2s - accuracy: 0.9947 - loss: 0.0205

193/938 [=====>........................] - ETA: 2s - accuracy: 0.9947 - loss: 0.0204

208/938 [=====>........................] - ETA: 2s - accuracy: 0.9946 - loss: 0.0208




































































































Learning rate for epoch 11 is 9.999999747378752e-06


Epoch 12/12


  1/938 [..............................] - ETA: 0s - accuracy: 1.0000 - loss: 0.0054

 16/938 [..............................] - ETA: 3s - accuracy: 0.9951 - loss: 0.0271

 31/938 [..............................] - ETA: 3s - accuracy: 0.9965 - loss: 0.0212

 46/938 [>.............................] - ETA: 3s - accuracy: 0.9952 - loss: 0.0217

 61/938 [>.............................] - ETA: 3s - accuracy: 0.9959 - loss: 0.0187

 76/938 [=>............................] - ETA: 2s - accuracy: 0.9963 - loss: 0.0174

 91/938 [=>............................] - ETA: 2s - accuracy: 0.9955 - loss: 0.0183

106/938 [==>...........................] - ETA: 2s - accuracy: 0.9959 - loss: 0.0171

121/938 [==>...........................] - ETA: 2s - accuracy: 0.9961 - loss: 0.0165

135/938 [===>..........................] - ETA: 2s - accuracy: 0.9959 - loss: 0.0165

149/938 [===>..........................] - ETA: 2s - accuracy: 0.9958 - loss: 0.0166

164/938 [====>.........................] - ETA: 2s - accuracy: 0.9953 - loss: 0.0178

179/938 [====>.........................] - ETA: 2s - accuracy: 0.9950 - loss: 0.0185

194/938 [=====>........................] - ETA: 2s - accuracy: 0.9949 - loss: 0.0185

209/938 [=====>........................] - ETA: 2s - accuracy: 0.9951 - loss: 0.0183




































































































Learning rate for epoch 12 is 9.999999747378752e-06


<tensorflow.python.keras.callbacks.History at 0x7f4bcc145860>

As you can see below, the checkpoints are getting saved.

In [16]:
# check the checkpoint directory
!ls {checkpoint_dir}

checkpoint		     ckpt_4.data-00000-of-00002
ckpt_1.data-00000-of-00002   ckpt_4.data-00001-of-00002
ckpt_1.data-00001-of-00002   ckpt_4.index
ckpt_1.index		     ckpt_5.data-00000-of-00002
ckpt_10.data-00000-of-00002  ckpt_5.data-00001-of-00002
ckpt_10.data-00001-of-00002  ckpt_5.index
ckpt_10.index		     ckpt_6.data-00000-of-00002
ckpt_11.data-00000-of-00002  ckpt_6.data-00001-of-00002
ckpt_11.data-00001-of-00002  ckpt_6.index
ckpt_11.index		     ckpt_7.data-00000-of-00002
ckpt_12.data-00000-of-00002  ckpt_7.data-00001-of-00002
ckpt_12.data-00001-of-00002  ckpt_7.index
ckpt_12.index		     ckpt_8.data-00000-of-00002
ckpt_2.data-00000-of-00002   ckpt_8.data-00001-of-00002
ckpt_2.data-00001-of-00002   ckpt_8.index
ckpt_2.index		     ckpt_9.data-00000-of-00002
ckpt_3.data-00000-of-00002   ckpt_9.data-00001-of-00002
ckpt_3.data-00001-of-00002   ckpt_9.index
ckpt_3.index


To see how the model perform, load the latest checkpoint and call `evaluate` on the test data.

Call `evaluate` as before using appropriate datasets.

In [17]:
model.load_weights(tf.train.latest_checkpoint(checkpoint_dir))

eval_loss, eval_acc = model.evaluate(eval_dataset)

print('Eval loss: {}, Eval Accuracy: {}'.format(eval_loss, eval_acc))

INFO:tensorflow:Reduce to /job:localhost/replica:0/task:0/device:CPU:0 then broadcast to ('/job:localhost/replica:0/task:0/device:CPU:0',).


INFO:tensorflow:Reduce to /job:localhost/replica:0/task:0/device:CPU:0 then broadcast to ('/job:localhost/replica:0/task:0/device:CPU:0',).


INFO:tensorflow:Reduce to /job:localhost/replica:0/task:0/device:CPU:0 then broadcast to ('/job:localhost/replica:0/task:0/device:CPU:0',).


INFO:tensorflow:Reduce to /job:localhost/replica:0/task:0/device:CPU:0 then broadcast to ('/job:localhost/replica:0/task:0/device:CPU:0',).


  1/157 [..............................] - ETA: 0s - accuracy: 0.9688 - loss: 0.1011

 10/157 [>.............................] - ETA: 0s - accuracy: 0.9828 - loss: 0.0480

 18/157 [==>...........................] - ETA: 0s - accuracy: 0.9861 - loss: 0.0409

 26/157 [===>..........................] - ETA: 0s - accuracy: 0.9862 - loss: 0.0374

 34/157 [=====>........................] - ETA: 0s - accuracy: 0.9881 - loss: 0.0376































Eval loss: 0.040106043219566345, Eval Accuracy: 0.9858999848365784


To see the output, you can download and view the TensorBoard logs at the terminal.

```
$ tensorboard --logdir=path/to/log-directory
```

In [18]:
!ls -sh ./logs

total 4.0K
4.0K train


## Export to SavedModel

Export the graph and the variables to the platform-agnostic SavedModel format. After your model is saved, you can load it with or without the scope.


In [19]:
path = 'saved_model/'

In [20]:
model.save(path, save_format='tf')

Instructions for updating:
If using Keras pass *_constraint arguments to layers.


Instructions for updating:
If using Keras pass *_constraint arguments to layers.


INFO:tensorflow:Assets written to: saved_model/assets


INFO:tensorflow:Assets written to: saved_model/assets


Load the model without `strategy.scope`.

In [21]:
unreplicated_model = tf.keras.models.load_model(path)

unreplicated_model.compile(
    loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True),
    optimizer=tf.keras.optimizers.Adam(),
    metrics=['accuracy'])

eval_loss, eval_acc = unreplicated_model.evaluate(eval_dataset)

print('Eval loss: {}, Eval Accuracy: {}'.format(eval_loss, eval_acc))

  1/157 [..............................] - ETA: 0s - loss: 0.1011 - accuracy: 0.9688

 12/157 [=>............................] - ETA: 0s - loss: 0.0456 - accuracy: 0.9844

 23/157 [===>..........................] - ETA: 0s - loss: 0.0383 - accuracy: 0.9857

 34/157 [=====>........................] - ETA: 0s - loss: 0.0376 - accuracy: 0.9881

























Eval loss: 0.040106043219566345, Eval Accuracy: 0.9858999848365784


Load the model with `strategy.scope`.

In [22]:
with strategy.scope():
  replicated_model = tf.keras.models.load_model(path)
  replicated_model.compile(loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True),
                           optimizer=tf.keras.optimizers.Adam(),
                           metrics=['accuracy'])

  eval_loss, eval_acc = replicated_model.evaluate(eval_dataset)
  print ('Eval loss: {}, Eval Accuracy: {}'.format(eval_loss, eval_acc))

  1/157 [..............................] - ETA: 0s - accuracy: 0.9688 - loss: 0.1011

 10/157 [>.............................] - ETA: 0s - accuracy: 0.9828 - loss: 0.0480

 19/157 [==>...........................] - ETA: 0s - accuracy: 0.9868 - loss: 0.0390

 28/157 [====>.........................] - ETA: 0s - accuracy: 0.9866 - loss: 0.0384































Eval loss: 0.040106043219566345, Eval Accuracy: 0.9858999848365784


### Examples and Tutorials
Here are some examples for using distribution strategy with keras fit/compile:
1. [Transformer](https://github.com/tensorflow/models/blob/master/official/nlp/transformer/transformer_main.py) example trained using `tf.distribute.MirroredStrategy`
2. [NCF](https://github.com/tensorflow/models/blob/master/official/recommendation/ncf_keras_main.py) example trained using `tf.distribute.MirroredStrategy`.

More examples listed in the [Distribution strategy guide](../../guide/distributed_training.ipynb#examples_and_tutorials)

## Next steps

* Read the [distribution strategy guide](../../guide/distributed_training.ipynb).
* Read the [Distributed Training with Custom Training Loops](training_loops.ipynb) tutorial.
* Visit the [Performance section](../../guide/function.ipynb) in the guide to learn more about other strategies and [tools](../../guide/profiler.md) you can use to optimize the performance of your TensorFlow models.

Note: `tf.distribute.Strategy` is actively under development and we will be adding more examples and tutorials in the near future. Please give it a try. We welcome your feedback via [issues on GitHub](https://github.com/tensorflow/tensorflow/issues/new).