Функция mixed_precision.Policy
Функция mixed_precision.Policy создает объект политики,
который управляет тем, какие типы данных используются
для вычислений и хранения переменных в слоях и моделях.
Первым параметром передается имя политики в виде строки.
Вторым необязательным параметром можно передать словарь
с настройками, например, указать тип данных для вычислений
и тип данных для переменных.
Политика смешанной точности позволяет ускорить обучение
и снизить потребление памяти за счет использования
float16 в вычислениях при сохранении
float32 для переменных модели.
Синтаксис
tf.keras.mixed_precision.Policy(name, [config])
Пример
Давайте создадим политику смешанной точности
с именем mixed_float16 и выведем ее параметры:
import tensorflow as tf
policy = tf.keras.mixed_precision.Policy('mixed_float16')
print(policy)
Результат выполнения кода:
"<Policy 'mixed_float16'>"
Пример
Давайте посмотрим на атрибуты политики, которые определяют типы данных для переменных и вычислений:
import tensorflow as tf
policy = tf.keras.mixed_precision.Policy('mixed_float16')
print(policy.variable_dtype)
print(policy.compute_dtype)
Результат выполнения кода:
"float32"
"float16"
Пример
Давайте установим созданную политику как глобальную и создадим слой Dense, чтобы увидеть, какой тип данных будет использоваться:
import tensorflow as tf
policy = tf.keras.mixed_precision.Policy('mixed_float16')
tf.keras.mixed_precision.set_global_policy(policy)
layer = tf.keras.layers.Dense(5)
t = tf.constant([[1, 2, 3], [4, 5, 6]])
res = layer(t)
print(res.dtype)
Результат выполнения кода:
"float16"
Пример
Давайте создадим политику с явно заданными типами данных через словарь конфигурации:
import tensorflow as tf
policy = tf.keras.mixed_precision.Policy(
'custom',
{'variable_dtype': 'float32', 'compute_dtype': 'bfloat16'}
)
print(policy.variable_dtype)
print(policy.compute_dtype)
Результат выполнения кода:
"float32"
"bfloat16"
Смотрите также
-
функцию
mixed_precision.set_global_policy,
которая устанавливает глобальную политику смешанной точности -
функцию
mixed_precision.global_policy,
которая возвращает текущую глобальную политику -
класс
LossScaleOptimizer,
который применяет масштабирование потерь при смешанной точности -
функцию
device,
которая задает устройство для выполнения операций