Spec-Zone.ru › TensorFlow

tf.while_loop

Повторяйте body, пока условие cond истинно. (устаревшие значения аргументов)

tf.while_loop(
    cond,
    body,
    loop_vars,
    shape_invariants=None,
    parallel_iterations=10,
    back_prop=True,
    swap_memory=False,
    maximum_iterations=None,
    name=None
)

Использование в блокнотах

Используется в руководстве Используется в учебных пособиях
  • Типы расширений
  • Моделирование распространения COVID-19 в Европе и воздействие мер вмешательства
  • Линейные смешанные эффекты моделей
Устарело: НЕКОТОРЫЕ ЗНАЧЕНИЯ АРГУМЕНТОВ УСТАРЕЛИ: (back_prop=False). Они будут удалены в будущей версии. Инструкции по обновлению: back_prop=False устарел. Рассмотрите использование tf.stop_gradient вместо него. Вместо этого: results = tf.while_loop(c, b, vars, back_prop=False) Используйте: results = tf.nest.map_structure(tf.stop_gradient, tf.while_loop(c, b, vars))
Примечание: Этот оператор автоматически используется в tf.function для преобразования циклов Python for и while, когда переменная цикла является tf.Tensor, если явно не указано autograph=False в tf.function аргументах. Например, следующие выражения эквивалентны:
@tf.function
def sumSquare(n):
  i, result = tf.constant(0), tf.constant(0)
  while i < n: # AutoGraph converts while-loop to tf.while_loop().
    result += i * i
    i += 1
  return result
sumSquare(10).numpy()
285
@tf.function
def sumSquare2(n):
  i, result = tf.constant(0), tf.constant(0)
  c = lambda i, _: tf.less(i, n)
  b = lambda i, result: (i + 1, result + i * i)
  return tf.while_loop(c, b, [i, result])[1]
sumSquare2(10).numpy()
285

Для получения дополнительной информации см. Руководство по tf.function и AutoGraph .

cond — это вызываемый объект, возвращающий скалярный булев тензор. body — это вызываемый объект, возвращающий (возможно, вложенную) кортеж, именованный кортеж или список тензоров той же размерности (длины и структуры) и типов, что и loop_vars. loop_vars — это (возможно, вложенная) кортеж, именованный кортеж или список тензоров, которые передаются как в cond, так и в body. cond и body принимают по столько же аргументов, сколько loop_vars.

В дополнение к обычным тензорам или IndexedSlices, тело может принимать и возвращать объекты TensorArray. Потоки объектов TensorArray будут соответствующим образом передаваться между циклами и во время вычисления градиентов.

Обратите внимание, что while_loop вызывает cond и body *точно один раз* (внутри вызова while_loop, и совсем не во время Session.run()). while_loop соединяет фрагменты графа, созданные во время вызовов cond и body, с дополнительными узлами графа, чтобы создать поток графа, который повторяет body, пока cond не вернёт ложь.

Для корректной работы tf.while_loop() строго обеспечивает инвариантность формы для переменных цикла. Инвариант формы — это (возможно, частичная) форма, которая не изменяется во время итераций цикла. Будет выдано сообщение об ошибке, если форма переменной цикла после итерации окажется более общей или несовместимой с её инвариантом формы. Например, форма [11, None] более общая, чем форма [11, 17], а форма [11, 21] несовместима с [11, 17]. По умолчанию (если аргумент shape_invariants не указан), предполагается, что начальная форма каждого тензора в loop_vars одинакова во всех итерациях. Аргумент shape_invariants позволяет вызывающему объекту указать менее точный инвариант формы для каждой переменной цикла, что необходимо, если форма меняется между итерациями. Функция tf.Tensor.set_shape также может использоваться в функции body, чтобы указать, что выходная переменная цикла имеет определённую форму.

Инвариант формы для SparseTensor и IndexedSlices обрабатывается следующим образом:

а) Если переменная цикла является SparseTensor, инвариант формы должен быть TensorShape([r]), где r — ранг плотного тензора, представленного разреженным тензором. Это означает, что формы трёх тензоров SparseTensor являются ([None], [None, r], [r]). ПРИМЕЧАНИЕ: Инвариант формы здесь — форма свойства SparseTensor.dense_shape. Это должна быть форма вектора.

б) Если переменная цикла является IndexedSlices, инвариант формы должен быть инвариантом формы тензора значений IndexedSlices. Это означает, что формы трёх тензоров IndexedSlices являются (shape, [shape[0]], [shape.ndims]).

while_loop реализует нестрогую семантику, позволяя выполнять несколько итераций параллельно. Максимальное количество параллельных итераций можно контролировать с помощью parallel_iterations, что даёт пользователям некоторый контроль над потреблением памяти и порядком выполнения. Для корректных программ while_loop должен возвращать один и тот же результат для любого parallel_iterations > 0.

Для обучения TensorFlow сохраняет тензоры, которые производятся в прямом выводе и необходимы в обратном распространении. Эти тензоры являются основным источником потребления памяти и часто вызывают ошибки OOM при обучении на графических процессорах. Когда флаг swap_memory имеет значение true, эти тензоры меняются с графического процессора на центральный процессор. Это, например, позволяет обучать модели RNN с очень длинными последовательностями и большими пакетами.

Аргументы
cond Вызываемый объект, представляющий условие завершения цикла.
body Вызываемый объект, представляющий тело цикла.
loop_vars (Возможная вложенная) кортеж, именованный кортеж или список numpy array, Tensor и TensorArray объектов.
shape_invariants Инварианты формы для переменных цикла.
parallel_iterations Количество итераций, разрешённых для выполнения параллельно. Должно быть положительным целым числом.
back_prop (Необязательно) Устарело. False отключает поддержку обратного распространения. Предпочтительно использовать tf.stop_gradient вместо этого.
swap_memory Включить обмен памятью между GPU и CPU для этого цикла.
maximum_iterations Необязательное максимальное количество итераций цикла while для выполнения. Если предоставлено, выход cond объединяется с дополнительным условием, гарантирующим, что количество выполненных итераций не превысит maximum_iterations.
name Необязательный префикс имени для возвращаемых тензоров.
Возвращаемое значение
Тензоры вывода для переменных цикла после цикла. Значение возврата имеет такую же структуру, что и loop_vars.
Исключения
TypeError если cond или body не являются вызываемыми объектами.
ValueError если loop_vars пусто.

Пример:

i = tf.constant(0)
c = lambda i: tf.less(i, 10)
b = lambda i: (tf.add(i, 1), )
r = tf.while_loop(c, b, [i])[0]
r.numpy()
10

Пример с вложенностью и именованным кортежем:

import collections
Pair = collections.namedtuple('Pair', 'j, k')
ijk_0 = (tf.constant(0), Pair(tf.constant(1), tf.constant(2)))
c = lambda i, p: i < 10
b = lambda i, p: (i + 1, Pair((p.j + p.k), (p.j - p.k)))
ijk_final = tf.while_loop(c, b, ijk_0)[1]
ijk_final[0].numpy(), ijk_final[1].numpy()
(32, 64)

Пример с использованием shape_invariants:

i0 = tf.constant(0)
m0 = tf.ones([2, 2])
c = lambda i, m: i < 10
b = lambda i, m: [i+1, tf.concat([m, m], axis=0)]
tf.while_loop(
    c, b, loop_vars=[i0, m0],
    shape_invariants=[i0.get_shape(), tf.TensorShape([None, 2])])[1]
<tf.Tensor: shape=(2048, 2), dtype=float32, numpy=...>

Пример, демонстрирующий нестрогую семантику: в следующем примере конечное значение counter не зависит от x. Таким образом, while_loop может увеличивать счётчик параллельно обновлениям x. Однако, поскольку счётчик цикла на одной итерации цикла зависит от значения на предыдущей итерации, сам счётчик цикла не может быть увеличен параллельно. Следовательно, если нам нужно только конечное значение счётчика (которое мы печатаем на строке print(sess.run(i))), то x никогда не будет увеличен, но счётчик будет обновлён на одном потоке. Напротив, если нам нужно значение вывода (которое мы печатаем на строке print(sess.run(out).shape)), то счётчик может быть увеличен в своём потоке, в то время как x может быть увеличен параллельно в отдельном потоке. В крайнем случае, можно предположить, что поток, увеличивающий счётчик, завершает работу, прежде чем x будет увеличен даже один раз. Единственное, что никогда не произойдёт, это то, что поток, обновляющий x, никогда не обгонит поток счётчика, поскольку поток, увеличивающий x, зависит от значения счётчика.

with tf.compat.v1.Session() as sess:
  n = 10
  c = lambda i, x: i < n
  b = lambda i, x: (
      tf.compat.v1.Print(i + 1, [i], "Updating i based on i == "),
      # Let x depend on i
      tf.compat.v1.Print(x + i, [i], "Updating x based on i == "))

  # Make x to be a big matrix so its updating thread would run slowly
  x = tf.zeros([1000, 100], dtype=tf.int32)
  counter = tf.constant(0)
  counter_out, x_out = tf.while_loop(c, b, (counter, x))

  # The following line may increment the counter and x in parallel.
  # The counter thread may get ahead of the x thread, but not the
  # other way around. For example, the log may contain these messages:
  #

... # Обновление i, основываясь на i == [9] ... # Обновление x, основываясь на i == [3] ... # ... # meaning that the counter(i) thread is on iteration 9, ... # while the x thread is on iteration 3. ... print(sess.run(x_out).shape) (1000, 100)

© 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/while_loop

Spec-Zone.ru

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