Spec-Zone.ru › TensorFlow 1.15

tf.while_loop

Просмотреть исходный код на GitHub

Повторять body до тех пор, пока условие cond истинно.

Просмотр псевдонимов

Псевдонимы для миграции

См. Руководство по миграции для получения дополнительных сведений.

tf.compat.v1.while_loop

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

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 не вернёт false.

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

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

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

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

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

Аргументы
cond Вызываемый объект, представляющий условие завершения цикла.
body Вызываемый объект, представляющий тело цикла.
loop_vars (Возможно, вложенный) кортеж, именованный кортеж или список numpy-массивов, Tensor, и TensorArray объектов.
shape_invariants Инварианты формы для переменных цикла.
parallel_iterations Количество итераций, разрешенных для параллельного выполнения. Должно быть положительным целым числом.
back_prop Включено ли обратное распространение для этого цикла while.
swap_memory Включен ли обмен памятью между графическим процессором и процессором для этого цикла.
name Необязательный префикс имени для возвращаемых тензоров.
maximum_iterations Необязательное максимальное число итераций цикла while, которое нужно выполнить. При указании результатом cond и результатом дополнительного условия, гарантирующего, что число выполненных итераций не превышает maximum_iterations.
return_same_structure Если True, вывод имеет ту же структуру, что и loop_vars. Если выполняется немедленное выполнение, это игнорируется (и всегда обрабатывается как True).
Возвращает
Выходные тензоры для переменных цикла после завершения цикла. Если return_same_structure равно True, значение возврата имеет ту же структуру, что и loop_vars. Если return_same_structure равно False, значение возврата — это тензор, TensorArray или IndexedSlice, если длина loop_vars равна 1, или список в противном случае.
Возбуждает исключения
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])

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

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)

Пример с использованием 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])])

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

import tensorflow as tf

n = 10000
x = tf.constant(list(range(n)))
c = lambda i, x: i < n
b = lambda i, x: (tf.compat.v1.Print(i + 1, [i]), tf.compat.v1.Print(x + 1,
[i], "x:"))
i, out = tf.while_loop(c, b, (0, x))
with tf.compat.v1.Session() as sess:
    print(sess.run(i))  # prints [0] ... [9999]

    # The following line may increment the counter and x in parallel.
    # The counter thread may get ahead of the other thread, but not the
    # other way around. So you may see things like
    # [9996] x:[9987]
    # meaning that the counter thread is on iteration 9996,
    # while the other thread is on iteration 9987
    print(sess.run(out).shape)

© 2020 The TensorFlow Authors. All rights reserved.
Licensed under the Creative Commons Attribution License 3.0.
Code samples licensed under the Apache 2.0 License.
https://www.tensorflow.org/versions/r1.15/api_docs/python/tf/while_loop

Spec-Zone.ru

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