tf.compat.v1.nn.raw_rnn
Создаёт RNN, заданный ячейкой RNN 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 — целое число, это должен быть Tensor соответствующего типа и формы [batch_size, cell.state_size]. Если cell.state_size — TensorShape, это должен быть Tensor соответствующего типа и формы [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 | Прозрачно меняет местами тензоры, созданные в прямом выводе, но необходимые для обратного распространения, с графического процессора на процессор. Это позволяет обучать RNN, которые, как правило, не помещаются на одном графическом процессоре, с минимальными (или без) потерями производительности. |
scope | Область переменных для созданного подграфа; по умолчанию "rnn". |
| Возвращаемые значения | |
|---|---|
Кортеж (emit_ta, final_state, final_loop_state), где:
|
| Исключения | |
|---|---|
TypeError | Если cell не является экземпляром RNNCell, или loop_fn не является callable. |
© 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/r2.3/api_docs/python/tf/compat/v1/nn/raw_rnn