Функция get_rng_state
Функция get_rng_state возвращает текущее состояние генератора
псевдослучайных чисел PyTorch на CPU. Состояние представлено в виде
тензора типа torch.ByteTensor. Это состояние можно сохранить,
а затем восстановить с помощью функции set_rng_state,
чтобы воспроизвести последовательности случайных чисел.
Функция не принимает никаких параметров.
Синтаксис
torch.get_rng_state()
Пример
Давайте получим текущее состояние генератора случайных чисел:
import torch
state = torch.get_rng_state()
print(state)
Результат выполнения кода:
tensor([ 0, 0, 0, ..., 0, 0, 0], dtype=torch.uint8)
Пример
Давайте сохраним состояние, сгенерируем несколько чисел, затем восстановим состояние и сгенерируем числа заново:
import torch
torch.manual_seed(0)
state = torch.get_rng_state()
t1 = torch.randint(0, 10, (3,))
print(t1)
torch.set_rng_state(state)
t2 = torch.randint(0, 10, (3,))
print(t2)
Результат выполнения кода:
tensor([4, 9, 3])
tensor([4, 9, 3])
Как видно из примера, последовательность случайных чисел полностью повторяется после восстановления состояния.
Пример
Состояние генератора можно сохранить в файл и загрузить позже:
import torch
torch.manual_seed(42)
state = torch.get_rng_state()
torch.save(state, 'rng_state.pt')
t = torch.rand(3)
print(t)
loaded_state = torch.load('rng_state.pt')
torch.set_rng_state(loaded_state)
t_restored = torch.rand(3)
print(t_restored)
Результат выполнения кода:
tensor([0.8823, 0.9150, 0.3829])
tensor([0.8823, 0.9150, 0.3829])
Смотрите также
-
функцию
set_rng_state,
которая восстанавливает состояние генератора случайных чисел -
функцию
manual_seed,
которая задаёт начальное зерно для генератора случайных чисел -
функцию
seed,
которая инициализирует генератор случайным зерном -
функцию
initial_seed,
которая возвращает начальное зерно генератора