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

Класс 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,
    который применяет линейное преобразование к входным данным
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить