Класс RNNCell
Класс RNNCell реализует один шаг простой рекуррентной нейронной сети.
Он принимает входной вектор и предыдущее скрытое состояние,
а возвращает новое скрытое состояние.
В отличие от RNN, этот класс не обрабатывает последовательность целиком,
а работает с одним временным шагом.
Это позволяет гибко управлять процессом обработки последовательностей.
Основные параметры конструктора:
-
input_size- размер входного вектора -
hidden_size- размер скрытого состояния -
bias- использовать ли смещение (по умолчаниюTrue) -
nonlinearity- функция активации:'tanh'или'relu'(по умолчанию'tanh')
Синтаксис
torch.nn.RNNCell(input_size, hidden_size, bias=True, nonlinearity='tanh')
Пример использования
Создадим ячейку RNN и применим её к одному временному шагу:
import torch
import torch.nn as nn
# Create RNN cell
rnn_cell = nn.RNNCell(input_size=3, hidden_size=5)
# Input tensor for one time step
x = torch.randn(2, 3) # batch=2, input_size=3
# Initial hidden state
h0 = torch.zeros(2, 5) # batch=2, hidden_size=5
# Forward pass
h1 = rnn_cell(x, h0)
print(h1.shape)
Результат выполнения кода:
torch.Size([2, 5])
Обработка последовательности
Для обработки всей последовательности в цикле передавайте скрытое состояние на каждом шаге:
import torch
import torch.nn as nn
rnn_cell = nn.RNNCell(input_size=3, hidden_size=5)
# Sequence of 4 time steps
x_seq = torch.randn(4, 2, 3) # seq_len=4, batch=2, input_size=3
h = torch.zeros(2, 5) # Initial hidden state
# Process sequence step by step
outputs = []
for x_t in x_seq:
h = rnn_cell(x_t, h)
outputs.append(h)
outputs = torch.stack(outputs)
print(outputs.shape)
Результат выполнения кода:
torch.Size([4, 2, 5])
Использование с ReLU
Можно использовать функцию активации ReLU вместо гиперболического тангенса:
import torch
import torch.nn as nn
rnn_cell = nn.RNNCell(input_size=3, hidden_size=5, nonlinearity='relu')
x = torch.randn(2, 3)
h0 = torch.zeros(2, 5)
h1 = rnn_cell(x, h0)
print(h1.shape)
Результат выполнения кода:
torch.Size([2, 5])
Без смещения
Создадим ячейку без параметра смещения:
import torch
import torch.nn as nn
rnn_cell = nn.RNNCell(input_size=3, hidden_size=5, bias=False)
x = torch.randn(2, 3)
h0 = torch.zeros(2, 5)
h1 = rnn_cell(x, h0)
print(h1.shape)
Результат выполнения кода:
torch.Size([2, 5])