tf.compat.v1.nn.raw_rnn
Создаёт RNN, заданный RNNCell cell и функцией цикла loop_fn.
tf.compat.v1.nn.raw_rnn(
cell, loop_fn, parallel_iterations=None, swap_memory=False, scope=None
)
Примечание: Этот метод всё ещё находится на стадии тестирования, и API может измениться.**
Эта функция является более примитивной версией dynamic_rnn, которая обеспечивает более прямой доступ к входным данным на каждой итерации. Она также обеспечивает больший контроль над началом и окончанием чтения последовательности и над тем, что выводить для вывода.
Например, она может использоваться для реализации динамического декодера модели seq2seq.
Вместо работы с Tensor объектами, большинство операций работают с TensorArray объектами напрямую.
Работа raw_rnn, в псевдокоде, в основном выглядит следующим образом:
time = tf.constant(0, dtype=tf.int32)
(finished, next_input, initial_state, emit_structure, loop_state) = loop_fn(
time=time, cell_output=None, cell_state=None, loop_state=None)
emit_ta = TensorArray(dynamic_size=True, dtype=initial_state.dtype)
state = initial_state
while not all(finished):
(output, cell_state) = cell(next_input, state)
(next_finished, next_input, next_state, emit, loop_state) = loop_fn(
time=time + 1, cell_output=output, cell_state=cell_state,
loop_state=loop_state)
# Emit zeros and copy forward state for minibatch entries that are finished.
state = tf.where(finished, state, next_state)
emit = tf.where(finished, tf.zeros_like(emit_structure), emit)
emit_ta = emit_ta.write(time, emit)
# If any new minibatch entries are marked as finished, mark these.
finished = tf.logical_or(finished, next_finished)
time += 1
return (emit_ta, state, loop_state)
с дополнительными свойствами, что выход и состояние могут быть (возможно, вложенными) кортежами, как определяется cell.output_size и cell.state_size, и в результате конечные state и emit_ta сами могут быть кортежами.
Простой пример реализации dynamic_rnn с помощью raw_rnn выглядит так:
inputs = tf.compat.v1.placeholder(shape=(max_time, batch_size, input_depth),
dtype=tf.float32)
sequence_length = tf.compat.v1.placeholder(shape=(batch_size,),
dtype=tf.int32)
inputs_ta = tf.TensorArray(dtype=tf.float32, size=max_time)
inputs_ta = inputs_ta.unstack(inputs)
cell = tf.compat.v1.nn.rnn_cell.LSTMCell(num_units)
def loop_fn(time, cell_output, cell_state, loop_state):
emit_output = cell_output # == None for time == 0
if cell_output is None: # time == 0
next_cell_state = cell.zero_state(batch_size, tf.float32)
else:
next_cell_state = cell_state
elements_finished = (time >= sequence_length)
finished = tf.reduce_all(elements_finished)
next_input = tf.cond(
finished,
lambda: tf.zeros([batch_size, input_depth], dtype=tf.float32),
lambda: inputs_ta.read(time))
next_loop_state = None
return (elements_finished, next_input, next_cell_state,
emit_output, next_loop_state)
outputs_ta, final_state, _ = raw_rnn(cell, loop_fn)
outputs = outputs_ta.stack()
| Аргументы | |
|---|---|
cell | Экземпляр RNNCell. |
loop_fn | Функция, принимающая входные данные (time, cell_output, cell_state, loop_state) и возвращающая кортеж (finished, next_input, next_cell_state, emit_output, next_loop_state). Здесь time — скаляр int32 Tensor, cell_output — Tensor или (возможно, вложенный) кортеж тензоров, как определено cell.output_size, и cell_state — Tensor или (возможно, вложенный) кортеж тензоров, как определено loop_fn при первом вызове (и должен соответствовать cell.state_size). Выходные данные: finished, логическое значение Tensor формы [batch_size], next_input: следующий вход для cell, next_cell_state: следующее состояние для cell, и emit_output: выходное значение для этой итерации. Обратите внимание, что emit_output должен быть Tensor или (возможно, вложенный) кортеж тензоров, который агрегируется в emit_ta внутри while_loop. При первом вызове loop_fn, значение emit_output соответствует значению emit_structure, которое затем используется для определения размера zero_tensor для emit_ta (по умолчанию cell.output_size). При последующих вызовах loop_fn, значение emit_output соответствует фактическому выходному тензору, который будет агрегирован в emit_ta. Параметр cell_state и выход next_cell_state могут быть либо одиночными, либо (возможно, вложенными) кортежами тензоров. Параметр loop_state и выход next_loop_state могут быть либо одиночными, либо (возможно, вложенными) кортежами объектов Tensor и TensorArray. Этот последний параметр может быть проигнорирован loop_fn, и возвращаемое значение может быть None. Если это не так None, то loop_state будет пропущен через цикл RNN, для использования исключительно loop_fn для отслеживания собственного состояния. Параметр next_loop_state может быть возвращен как None. Первый вызов loop_fn будет time = 0, cell_output = None, cell_state = None, и loop_state = None. Для этого вызова: значение next_cell_state должно быть значением, с которым инициализируется состояние ячейки. Это может быть конечное состояние из предыдущего RNN или результат вызова cell.zero_state(). Это должно быть (возможно, вложенное) кортежное структура тензоров. Если cell.state_size — целое число, это должен быть объект соответствующего типа и формы [batch_size, cell.state_size]. Если cell.state_size — TensorShape, это должен быть объект соответствующего типа и формы [batch_size] + cell.state_size. Если cell.state_size — (возможно, вложенный) кортеж целых чисел или TensorShape, это будет кортеж с соответствующими формами. Значение emit_output может быть либо None, либо (возможно, вложенная) кортежная структура тензоров, например, (tf.zeros(shape_0, dtype=dtype_0), tf.zeros(shape_1, dtype=dtype_1)). Если это первое значение emit_output возвращается как None, то результат emit_ta от raw_rnn будет иметь ту же структуру и типы данных, что и cell.output_size. Иначе emit_ta будет иметь ту же структуру, формы (префиксная размерность batch_size) и типы данных, что и emit_output. Фактические значения, возвращаемые emit_output при этом вызове инициализации, игнорируются. Обратите внимание, что эта структура вывода должна быть согласованной на всех временных шагах. |
parallel_iterations | (По умолчанию: 32). Количество итераций, выполняемых параллельно. Те операции, которые не имеют временной зависимости и могут выполняться параллельно, будут. Этот параметр позволяет балансировать время и память. Значения >> 1 требуют больше памяти, но занимают меньше времени, в то время как меньшие значения требуют меньше памяти, но вычисления занимают больше времени. |
swap_memory | Прозрачный обмен тензорами, полученными в процессе прямого вывода, но необходимыми для обратного распространения, между графическим процессором (GPU) и центральным процессором (CPU). Это позволяет обучать RNN, которые обычно не помещаются на одном GPU, с минимальными (или без) потерями производительности. |
scope | Область переменных (VariableScope) для созданной подграфа; по умолчанию "rnn". |
| Возвращаемое значение | |
|---|---|
Кортеж (emit_ta, final_state, final_loop_state), где:
|
| Исключения | |
|---|---|
TypeError | Если cell не является экземпляром RNNCell, или loop_fn не является callable. |
© 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/nn/raw_rnn