GRU
-
class torch.nn.GRU(*args, **kwargs)[source] -
Применяет многослойный рекуррентный блок GRU (gated recurrent unit) к последовательности входных данных.
Для каждого элемента в последовательности входных данных каждый слой вычисляет следующую функцию:
где — скрытое состояние во времени
t, — вход во времениt, — скрытое состояние слоя во времениt-1или начальное скрытое состояние во времени0, а , , — это, соответственно, ворота сброса, обновления и нового состояния. — сигмоидальная функция, а — произведение Адамара.В многослойном блоке GRU вход -го слоя () — это скрытое состояние предыдущего слоя, умноженного на значение дропаута , где каждое — случайная величина Бернулли, которая равна с вероятностью
dropout.- Параметры:
-
-
input_size – Количество ожидаемых признаков в входных данных
x -
hidden_size – Количество признаков в скрытом состоянии
h -
num_layers – Количество рекуррентных слоёв. Например, установка
num_layers=2означает объединение двух блоков GRU для формированияstacked GRU, где второй блок GRU принимает на вход выходные данные первого блока GRU и вычисляет конечные результаты. По умолчанию: 1 -
bias – Если
False, то слой не использует весы смещенияb_ihиb_hh. По умолчанию:True -
batch_first – Если
True, то тензоры входных и выходных данных предоставляются как(batch, seq, feature)вместо(seq, batch, feature). Обратите внимание, что это не относится к скрытым или ячейкам состояния. Подробности см. в разделах «Входы/Выходы» ниже. По умолчанию:False -
dropout – Если не равно нулю, вводит слой
Dropoutна выходах каждого слоя GRU, кроме последнего слоя, с вероятностью дропаута, равнойdropout. По умолчанию: 0 -
bidirectional – Если
True, становится двунаправленным блоком GRU. По умолчанию:False
-
input_size – Количество ожидаемых признаков в входных данных
- Входы: input, h_0
-
-
input: тензор формы для неразмеченного ввода, , когда
batch_first=Falseили , когдаbatch_first=True, содержащий признаки входной последовательности. Вход также может быть упакованной последовательностью переменной длины. Смотритеtorch.nn.utils.rnn.pack_padded_sequence()илиtorch.nn.utils.rnn.pack_sequence()для подробностей. - h_0: тензор формы или , содержащий начальное скрытое состояние для входной последовательности. По умолчанию равен нулю, если не указано.
где:
-
input: тензор формы для неразмеченного ввода, , когда
- Выходы: output, h_n
-
-
output: тензор формы для неразмеченного ввода, , когда
batch_first=Falseили когдаbatch_first=Trueсодержащий выходные признаки(h_t)из последнего слоя GRU для каждойt. Если в качестве входного значения предоставленаtorch.nn.utils.rnn.PackedSequence, выход также будет упакованной последовательностью. - h_n: тензор формы или содержащий конечное скрытое состояние для входной последовательности.
-
output: тензор формы для неразмеченного ввода, , когда
- Переменные:
-
-
weight_ih_l[k] – обученные веса вход-скрытое для слоя (W_ir|W_iz|W_in), формы
(3*hidden_size, input_size)дляk = 0. Иначе, форма(3*hidden_size, num_directions * hidden_size) -
weight_hh_l[k] – обученные веса скрытое-скрытое для слоя (W_hr|W_hz|W_hn), формы
(3*hidden_size, hidden_size) -
bias_ih_l[k] – обученный смещение вход-скрытое для слоя (b_ir|b_iz|b_in), формы
(3*hidden_size) -
bias_hh_l[k] – обученное смещение скрытое-скрытое для слоя (b_hr|b_hz|b_hn), формы
(3*hidden_size)
-
weight_ih_l[k] – обученные веса вход-скрытое для слоя (W_ir|W_iz|W_in), формы
Примечание
Все веса и смещения инициализируются из , где
Примечание
Для двунаправленных GRU, направление вперед и назад – 0 и 1 соответственно. Пример разделения выходных слоев, когда
batch_first=False:output.view(seq_len, batch, num_directions, hidden_size).Примечание
Аргумент
batch_firstигнорируется для неразмеченных входных данных.Примечание
Если выполняются следующие условия: 1) cudnn включен, 2) данные ввода находятся на графическом процессоре, 3) тип данных входных данных
torch.float16, 4) используется графический процессор V100, 5) входные данные не в форматеPackedSequence, может быть выбран алгоритм сохранения для повышения производительности.Примеры:
>>> rnn = nn.GRU(10, 20, 2) >>> input = torch.randn(5, 3, 10) >>> h0 = torch.randn(2, 3, 20) >>> output, hn = rnn(input, h0)
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/1.13/generated/torch.nn.GRU.html