Spec-Zone.ru › TensorFlow 2.9

tf.compat.v1.train.get_global_step

Получение тензора глобальной шаги.

tf.compat.v1.train.get_global_step(
    graph=None
)

Миграция на TF2

Внимание: Этот API был разработан для TensorFlow v1. Продолжайте чтение, чтобы узнать, как мигрировать из этого API в эквивалент в чистом TensorFlow v2. См. Руководство по миграции TensorFlow v1 в TensorFlow v2 для получения инструкций по миграции остальной части вашего кода.

В связи с устареванием глобальных графов, TF больше не отслеживает переменные в коллекциях. Другими словами, в TF2 нет глобальных переменных. Таким образом, функции глобальной шаги были удалены (get_or_create_global_step, create_global_step, get_global_step) . У вас есть два варианта миграции:

  1. Создайте оптимизатор Keras, который генерирует переменную iterations. Эта переменная автоматически увеличивается при вызове apply_gradients.
  2. Создайте и увеличьте переменную tf.Variable вручную.

Ниже приведен пример миграции от использования глобальной шаги к использованию оптимизатора Keras:

Определите модель и функцию потерь:

def compute_loss(x):
  v = tf.Variable(3.0)
  y = x * v
  loss = x * 5 - x * v
  return loss, [v]

До миграции:

g = tf.Graph()
with g.as_default():
  x = tf.compat.v1.placeholder(tf.float32, [])
  loss, var_list = compute_loss(x)
  global_step = tf.compat.v1.train.get_or_create_global_step()
  global_init = tf.compat.v1.global_variables_initializer()
  optimizer = tf.compat.v1.train.GradientDescentOptimizer(0.1)
  train_op = optimizer.minimize(loss, global_step, var_list)
sess = tf.compat.v1.Session(graph=g)
sess.run(global_init)
print("before training:", sess.run(global_step))
before training: 0
sess.run(train_op, feed_dict={x: 3})
print("after training:", sess.run(global_step))
after training: 1

Используя get_global_step:

with g.as_default():
  print(sess.run(tf.compat.v1.train.get_global_step()))
1

Миграция к оптимизатору Keras:

optimizer = tf.keras.optimizers.SGD(.01)
print("before training:", optimizer.iterations.numpy())
before training: 0
with tf.GradientTape() as tape:
  loss, var_list = compute_loss(3)
  grads = tape.gradient(loss, var_list)
  optimizer.apply_gradients(zip(grads, var_list))
print("after training:", optimizer.iterations.numpy())
after training: 1

Описание

Тензор глобальной шаги должен быть целочисленной переменной. Сначала мы пытаемся найти его в коллекции GLOBAL_STEP, или по имени global_step:0.

Аргументы
graph Граф для поиска глобальной шаги. Если отсутствует, используется стандартный граф.
Возвращаемые значения
Переменная глобальной шаги или None, если не найдена.
Исключения
TypeError Если тензор глобальной шаги имеет нецелочисленный тип или это не Variable.

© 2022 The TensorFlow Authors. All rights reserved.
Licensed under the Creative Commons Attribution License 4.0.
Code samples licensed under the Apache 2.0 License.
https://www.tensorflow.org/versions/r2.9/api_docs/python/tf/compat/v1/train/get_global_step

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API