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

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

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

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