tf.nn.RNNCellDropoutWrapper
Оператор добавления дропаута к входным и выходным данным заданного ячейки.
Наследуется от: AbstractRNNCell, Layer, Module
tf.nn.RNNCellDropoutWrapper(
*args, **kwargs
)
| Аргументы | |
|---|---|
cell | RNNCell, к которому добавляется проекция на размер вывода. |
input_keep_prob | тензор или число с плавающей точкой от 0 до 1, вероятность сохранения входных данных; если оно постоянно и равно 1, входной дропаут не будет добавлен. |
output_keep_prob | тензор или число с плавающей точкой от 0 до 1, вероятность сохранения выходных данных; если оно постоянно и равно 1, выходной дропаут не будет добавлен. |
state_keep_prob | тензор или число с плавающей точкой от 0 до 1, вероятность сохранения выходных данных; если оно постоянно и равно 1, выходной дропаут не будет добавлен. Дропаут состояния применяется к выходящим состояниям ячейки. Примечание, компоненты состояния, к которым применяется дропаут, когда state_keep_prob находится в (0, 1), также определяются аргументом dropout_state_filter_visitor (например, по умолчанию дропаут никогда не применяется к компоненту c объекта LSTMStateTuple). |
variational_recurrent | Python bool. Если True, то один и тот же шаблон дропаута применяется ко всем временным шагам при вызове. Если этот параметр установлен, input_size должен быть указан. |
input_size | (необязательно) (возможно, вложенная кортеж) объекты TensorShape, содержащие глубину(ы) входных тензоров, ожидаемых для передачи в DropoutWrapper. Требуется и используется если variational_recurrent = True и input_keep_prob < 1. |
dtype | (необязательно) Размерность входных, состояниевых и выходных тензоров. Требуется и используется если variational_recurrent = True. |
seed | (необязательно) целое число, семя генерации случайных чисел. |
dropout_state_filter_visitor | (необязательно), по умолчанию: (см. ниже). Функция, которая принимает любой иерархический уровень состояния и возвращает скаляр или структуру с глубиной 1 Python-булевых значений, описывающих, какие термины в состоянии должны быть исключены. Кроме того, если функция возвращает True, дропаут применяется на этом подуровне. Если функция возвращает False, дропаут не применяется на этом подуровне. Поведение по умолчанию: выполнять дропаут для всех терминов, кроме памяти (c) состояния объектов LSTMCellState, и не пытаться применить дропаут к объектам TensorArray: def dropout_state_filter_visitor(s): if isinstance(s, LSTMCellState): # Never perform dropout on the c state. return LSTMCellState(c=False, h=True) elif isinstance(s, TensorArray): return False return True |
**kwargs | словарь ключевых аргументов для базового слоя. |
| Исключения | |
|---|---|
TypeError | если cell не является RNNCell, или keep_state_fn указан, но не callable. |
ValueError | если какие-либо из keep_probs не находятся в диапазоне от 0 до 1. |
| Атрибуты | |
|---|---|
output_size | |
state_size | |
wrapped_cell | |
Методы
get_initial_state
get_initial_state(
inputs=None, batch_size=None, dtype=None
)
zero_state
zero_state(
batch_size, dtype
)
© 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.4/api_docs/python/tf/nn/RNNCellDropoutWrapper