РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
795 of 824 menu

Функция 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
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить