Класс LSTMCell
Класс LSTMCell реализует отдельную ячейку LSTM (Long Short-Term Memory) для пошаговой обработки последовательностей. В отличие от слоя LSTM, который обрабатывает всю последовательность целиком, LSTMCell принимает один временной шаг и возвращает новое скрытое состояние и новое состояние ячейки. Это позволяет управлять процессом обработки в цикле и реализовывать собственные архитектуры.
Параметры конструктора:
-
input_size- размерность входных данных на каждом временном шаге -
hidden_size- размерность скрытого состояния -
bias- еслиTrue(по умолчанию), то используются смещения -
device- устройство для размещения параметров (опционально) -
dtype- тип данных параметров (опционально)
Синтаксис
torch.nn.LSTMCell(input_size, hidden_size, bias=True, device=None, dtype=None)
Пример
Давайте создадим ячейку LSTM с размерностью входных данных 3 и размерностью скрытого состояния 5:
import torch
cell = torch.nn.LSTMCell(input_size=3, hidden_size=5)
print(cell)
Результат выполнения кода:
LSTMCell(3, 5)
Пример
Передадим в ячейку входной вектор размерности 3 и начальные состояния (скрытое и ячейки) размерности 5:
import torch
torch.manual_seed(0)
cell = torch.nn.LSTMCell(3, 5)
x = torch.randn(2, 3) # batch_size=2, input_size=3
h = torch.randn(2, 5) # batch_size=2, hidden_size=5
c = torch.randn(2, 5) # batch_size=2, hidden_size=5
h_new, c_new = cell(x, (h, c))
print(h_new.shape)
print(c_new.shape)
Результат выполнения кода:
torch.Size([2, 5])
torch.Size([2, 5])
Пример
Используем ячейку LSTM в цикле для обработки последовательности из трёх временных шагов:
import torch
torch.manual_seed(0)
cell = torch.nn.LSTMCell(3, 5)
seq = torch.randn(3, 2, 3) # sequence_length=3, batch_size=2, input_size=3
h = torch.zeros(2, 5)
c = torch.zeros(2, 5)
outputs = []
for t in range(seq.size(0)):
h, c = cell(seq[t], (h, c))
outputs.append(h)
res = torch.stack(outputs)
print(res.shape)
Результат выполнения кода:
torch.Size([3, 2, 5])
Пример
Создадим ячейку LSTM без смещений:
import torch
cell = torch.nn.LSTMCell(4, 6, bias=False)
print(cell)
Результат выполнения кода:
LSTMCell(4, 6, bias=False)