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