tf.compat.v1.nn.static_rnn
Создаёт рекуррентную нейронную сеть, определяемую объектом RNNCell cell. (устарело)
tf.compat.v1.nn.static_rnn(
cell, inputs, initial_state=None, dtype=None, sequence_length=None, scope=None
)
Простейшая форма генерируемой RNN-сети:
state = cell.zero_state(...) outputs = [] for input_ in inputs: output, state = cell(input_, state) outputs.append(output) return (outputs, state)
Однако доступно несколько других вариантов:
Можно указать начальное состояние. Если задан вектор sequence_length, выполняется динамическое вычисление. Этот метод вычисления не вычисляет шаги RNN, выходящие за пределы максимальной длины последовательности мини-пакета (что экономит вычислительное время) и правильно распространяет состояние в момент длины последовательности примера до конечного выходного состояния.
Выполняемое динамическое вычисление на момент t для строки мини-пакета b,
(output, state)(b, t) =
(t >= sequence_length(b))
? (zeros(cell.output_size), states(b, sequence_length(b) - 1))
: cell(input(b, t), state(b, t - 1))
| Аргументы | |
|---|---|
cell | Экземпляр RNNCell. |
inputs | Список из T входных значений, каждый из которых представляет собой Tensor формы [batch_size, input_size], или вложенную кортеж таких элементов. |
initial_state | (необязательно) Начальное состояние RNN. Если cell.state_size является целым числом, это должно быть Tensor соответствующего типа и формы [batch_size, cell.state_size]. Если cell.state_size является кортежем, это должен быть кортеж тензоров с формами [batch_size, s] for s in cell.state_size. |
dtype | (необязательно) Тип данных для начального состояния и ожидаемого выхода. Требуется, если initial_state не указан или состояние RNN имеет разнородный тип данных. |
sequence_length | Указывает длину каждой последовательности во входных данных. Вектор (тензор) int32 или int64 размера [batch_size], значения в [0, T). |
scope | Область переменных для созданного подграфа; по умолчанию "rnn". |
| Возвращаемое значение | |
|---|---|
Пара (выходы, состояние), где:
|
| Исключения | |
|---|---|
TypeError | Если cell не является экземпляром RNNCell. |
ValueError | Если inputs является None или пустым списком, или если глубина входа (размер столбца) не может быть определена из входных данных с помощью вывода формы. |
© 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/static_rnn