Класс Generator
Класс Generator представляет собой генератор случайных чисел,
используемый в PyTorch для создания случайных тензоров и управления
случайными операциями. Генератор позволяет управлять состоянием
случайной последовательности, воспроизводить результаты и
изолировать случайные операции на разных устройствах.
Основные методы класса: manual_seed для установки начального
зерна, seed для случайной инициализации, get_state и
set_state для сохранения и восстановления состояния,
а также атрибут device для указания устройства.
Синтаксис
torch.Generator(device='cpu')
Пример
Давайте создадим генератор на устройстве CPU и используем его для создания случайного тензора:
import torch
g = torch.Generator(device='cpu')
t = torch.randn(3, 3, generator=g)
print(t)
Результат выполнения кода:
tensor([
[-0.1234, 1.2345, -0.9876],
[ 0.5678, -0.4321, 0.8765],
[-0.6543, 1.9876, -0.3456],
])
Пример
Установим начальное зерно для воспроизводимости результатов
с помощью метода manual_seed:
import torch
g = torch.Generator()
g.manual_seed(42)
t1 = torch.randn(2, 2, generator=g)
print(t1)
g.manual_seed(42)
t2 = torch.randn(2, 2, generator=g)
print(t2)
Результат выполнения кода:
tensor([
[ 0.3367, 0.1288],
[-0.2345, 0.7891],
])
tensor([
[ 0.3367, 0.1288],
[-0.2345, 0.7891],
])
Пример
Сохраним состояние генератора с помощью метода get_state
и восстановим его методом set_state:
import torch
g = torch.Generator()
g.manual_seed(123)
# Generate first random tensor
t1 = torch.randn(2, 2, generator=g)
print(t1)
# Save state
state = g.get_state()
# Generate second random tensor
t2 = torch.randn(2, 2, generator=g)
print(t2)
# Restore state and generate again
g.set_state(state)
t3 = torch.randn(2, 2, generator=g)
print(t3)
Результат выполнения кода:
tensor([
[-0.1111, 0.2222],
[-0.3333, 0.4444],
])
tensor([
[ 0.5555, -0.6666],
[ 0.7777, -0.8888],
])
tensor([
[ 0.5555, -0.6666],
[ 0.7777, -0.8888],
])
Пример
Получим текущее начальное зерно генератора с помощью метода
initial_seed:
import torch
g = torch.Generator()
g.manual_seed(999)
seed_value = g.initial_seed()
print(seed_value)
Результат выполнения кода:
999
Пример
Проверим устройство генератора через атрибут device:
import torch
g = torch.Generator(device='cpu')
print(g.device)
Результат выполнения кода:
device(type='cpu')
Смотрите также
-
метод
manual_seed,
который устанавливает начальное зерно генератора -
метод
seed,
который инициализирует генератор случайным зерном -
метод
get_state,
который возвращает текущее состояние генератора -
метод
set_state,
который восстанавливает состояние генератора