RNNBase
-
class torch.nn.RNNBase(mode, input_size, hidden_size, num_layers=1, bias=True, batch_first=False, dropout=0.0, bidirectional=False, proj_size=0, device=None, dtype=None)[source] -
Базовый класс для модулей RNN (RNN, LSTM, GRU).
Реализует аспекты RNN, общие для классов RNN, LSTM и GRU, такие как инициализация модуля и вспомогательные методы для управления хранением параметров.
Примечание
Метод forward не реализован классом RNNBase.
Примечание
Классы LSTM и GRU переопределяют некоторые методы, реализованные в RNNBase.
-
flatten_parameters()[source] -
Сбрасывает указатель на данные параметров, чтобы они могли использовать более быстрые пути кода.
В настоящее время это работает только если модуль находится на GPU и включен cuDNN. В противном случае это бесполезная операция.
-
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.nn.RNNBase.html