Spec-Zone.ru › TensorFlow 2.4

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), где:

emit_ta: Вывод RNN TensorArray. Если loop_fn возвращает (возможно, вложенный) набор тензоров для emit_output во время инициализации (входы time = 0, cell_output = None и loop_state = None), то emit_ta будет иметь ту же структуру, типы данных и формы, что и emit_output вместо этого. Если loop_fn возвращает emit_output = None при этом вызове, то структура cell.output_size используется: Если cell.output_size является (возможно, вложенным) кортежем целых чисел или TensorShape объектов, то emit_ta будет кортежем, имеющим ту же структуру, что и cell.output_size, содержащим TensorArrays, чьи формы элементов соответствуют данным формы в cell.output_size.

final_state: Конечное состояние ячейки. Если cell.state_size — целое число, оно будет иметь форму [batch_size, cell.state_size]. Если это TensorShape, оно будет иметь форму [batch_size] + cell.state_size. Если это (возможно, вложенный) кортеж целых чисел или TensorShape, это будет кортеж с соответствующими формами.

final_loop_state: Конечное состояние цикла, возвращаемое loop_fn.

Исключения
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

Spec-Zone.ru

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