Класс GRUCell
Класс GRUCell представляет собой слой,
реализующий одну ячейку Gated Recurrent Unit
(GRU). В отличие от полноценного слоя GRU,
который обрабатывает всю последовательность,
ячейка работает с одним шагом по времени.
Первый параметр input_size определяет
размер входного вектора, второй
hidden_size задаёт размер скрытого
состояния. Класс поддерживает возможность
использования смещения (bias) и
различные варианты инициализации весов
(через параметр device и
dtype).
Синтаксис
torch.nn.GRUCell(
input_size,
hidden_size,
bias=True,
device=None,
dtype=None
)
Параметры
Метод принимает следующие параметры:
-
input_size- размерность входного вектора (количество признаков на одном временном шаге) -
hidden_size- размерность скрытого состояния (количество нейронов в ячейке) -
bias- флаг использования смещения (по умолчаниюTrue) -
device- устройство для размещения параметров (опционально) -
dtype- тип данных параметров (опционально)
Пример
Давайте создадим ячейку GRU с размером входа
10 и размером скрытого состояния
20:
import torch
gru_cell = torch.nn.GRUCell(
input_size=10,
hidden_size=20
)
print(gru_cell)
Результат выполнения кода:
GRUCell(10, 20)
Пример
Теперь применим ячейку к входному вектору
размера 10, передав начальное скрытое
состояние:
import torch
torch.manual_seed(0)
gru_cell = torch.nn.GRUCell(10, 20)
x = torch.randn(3, 10)
h0 = torch.randn(3, 20)
h1 = gru_cell(x, h0)
print(h1.shape)
Результат выполнения кода:
torch.Size([3, 20])
Пример
Если начальное состояние не передавать, оно автоматически заполняется нулями:
import torch
torch.manual_seed(0)
gru_cell = torch.nn.GRUCell(5, 8)
x = torch.randn(2, 5)
h1 = gru_cell(x)
print(h1)
Результат выполнения кода:
tensor([
[ 0.1724, -0.0041, 0.1055, -0.0778, -0.1064, -0.2783, -0.1641, 0.1104],
[ 0.0559, -0.0410, 0.1751, 0.0604, -0.1726, -0.1799, -0.1334, -0.0826]
], grad_fn=<AddBackward0>)
Пример
Покажем обработку последовательности из нескольких шагов, обновляя скрытое состояние на каждом шаге:
import torch
torch.manual_seed(0)
gru_cell = torch.nn.GRUCell(3, 5)
seq = torch.randn(4, 2, 3)
h = torch.zeros(2, 5)
outputs = []
for t in range(seq.size(0)):
h = gru_cell(seq[t], h)
outputs.append(h)
res = torch.stack(outputs)
print(res.shape)
Результат выполнения кода:
torch.Size([4, 2, 5])
Пример
Отключим использование смещения, создав
ячейку с параметром bias равным
False:
import torch
gru_cell = torch.nn.GRUCell(
input_size=4,
hidden_size=6,
bias=False
)
print(gru_cell.bias_ih is None)
print(gru_cell.bias_hh is None)
Результат выполнения кода:
True
True
Смотрите также
-
класс
GRU,
который реализует полноценный слой GRU для обработки всей последовательности -
класс
LSTMCell,
который реализует ячейку LSTM для одного временного шага -
класс
RNNCell,
который реализует простую ячейку RNN для одного временного шага -
класс
Linear,
который применяет линейное преобразование к входным данным