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/api_docs/python/tf/compat/v1/train/get_global_step