tf.compat.v1.train.get_global_step
Получение тензора глобальной шаги.
tf.compat.v1.train.get_global_step(
graph=None
)
Миграция на TF2
В связи с устареванием глобальных графов, TF больше не отслеживает переменные в коллекциях. Другими словами, в TF2 нет глобальных переменных. Таким образом, функции глобальной шаги были удалены (get_or_create_global_step, create_global_step, get_global_step) . У вас есть два варианта миграции:
- Создайте оптимизатор Keras, который генерирует переменную
iterations. Эта переменная автоматически увеличивается при вызовеapply_gradients. - Создайте и увеличьте переменную
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