РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
355 of 769 menu

Класс 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)

Смотрите также

  • класс LSTM,
    который обрабатывает всю последовательность целиком
  • класс GRUCell,
    который реализует ячейку Gated Recurrent Unit
  • класс RNNCell,
    который реализует простую рекуррентную ячейку
  • класс Linear,
    который применяет линейное преобразование к входным данным
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить