tf.keras.utils.multi_gpu_model
| Просмотреть исходный код на GitHub |
Реплицирует модель на разных графических процессорах.
tf.keras.utils.multi_gpu_model(
model, gpus, cpu_merge=True, cpu_relocation=False
)
В частности, эта функция реализует параллелизм данных с использованием нескольких графических процессоров на одной машине. Она работает следующим образом:
- Разделяет входные данные модели на несколько под-парт.
- Применяет копию модели к каждой под-парт. Каждая копия модели выполняется на выделенном графическом процессоре.
- Объединяет результаты (на процессоре) в одну большую партию.
Например, если ваш batch_size составляет 64, и вы используете gpus=2, то мы разделим вход на 2 под-парт по 32 образца, обработаем каждую под-парт на одном графическом процессоре, а затем вернём полную партию из 64 обработанных образцов.
Это обеспечивает квазилинейное ускорение до 8 графических процессоров.
В настоящее время эта функция доступна только с бэкендом TensorFlow.
| Аргументы | |
|---|---|
model | Экземпляр модели Keras. Чтобы избежать ошибок OOM, эта модель могла быть построена на процессоре, например (см. пример использования ниже). |
gpus | Целое число >= 2, количество графических процессоров, на которых необходимо создать реплики модели. |
cpu_merge | Булево значение, определяющее, нужно ли принудительно объединять веса модели в рамках области процессора. |
cpu_relocation | Булево значение, определяющее, нужно ли создавать веса модели в рамках области процессора. Если модель не определена в рамках какой-либо предыдущей области устройства, вы можете всё равно её восстановить, активировав этот параметр. |
| Возвращает | |
|---|---|
Экземпляр Keras Model , который может использоваться так же, как и исходный аргумент model, но который распределяет свою работу на нескольких графических процессорах. |
Пример 1: Обучение моделей с объединением весов на процессоре
import tensorflow as tf
from keras.applications import Xception
from keras.utils import multi_gpu_model
import numpy as np
num_samples = 1000
height = 224
width = 224
num_classes = 1000
# Instantiate the base model (or "template" model).
# We recommend doing this with under a CPU device scope,
# so that the model's weights are hosted on CPU memory.
# Otherwise they may end up hosted on a GPU, which would
# complicate weight sharing.
with tf.device('/cpu:0'):
model = Xception(weights=None,
input_shape=(height, width, 3),
classes=num_classes)
# Replicates the model on 8 GPUs.
# This assumes that your machine has 8 available GPUs.
parallel_model = multi_gpu_model(model, gpus=8)
parallel_model.compile(loss='categorical_crossentropy',
optimizer='rmsprop')
# Generate dummy data.
x = np.random.random((num_samples, height, width, 3))
y = np.random.random((num_samples, num_classes))
# This `fit` call will be distributed on 8 GPUs.
# Since the batch size is 256, each GPU will process 32 samples.
parallel_model.fit(x, y, epochs=20, batch_size=256)
# Save model via the template model (which shares the same weights):
model.save('my_model.h5')
Пример 2: Обучение моделей с объединением весов на процессоре с использованием cpu_relocation
..
# Not needed to change the device scope for model definition:
model = Xception(weights=None, ..)
try:
model = multi_gpu_model(model, cpu_relocation=True)
print("Training using multiple GPUs..")
except:
print("Training using single GPU or CPU..")
model.compile(..)
..
Пример 3: Обучение моделей с объединением весов на графическом процессоре (рекомендуется для NV-link)
..
# Not needed to change the device scope for model definition:
model = Xception(weights=None, ..)
try:
model = multi_gpu_model(model, cpu_merge=False)
print("Training using multiple GPUs..")
except:
print("Training using single GPU or CPU..")
model.compile(..)
..
| Возможные исключения | |
|---|---|
ValueError | если аргумент gpus не соответствует доступным устройствам. |
© 2020 The TensorFlow Authors. All rights reserved.
Licensed under the Creative Commons Attribution License 3.0.
Code samples licensed under the Apache 2.0 License.
https://www.tensorflow.org/versions/r1.15/api_docs/python/tf/keras/utils/multi_gpu_model