Функция cuda.set_rng_state
Функция cuda.set_rng_state восстанавливает состояние
генератора случайных чисел (ГСЧ) для конкретного устройства
CUDA. Первым параметром передаётся состояние в виде объекта
torch.ByteTensor, полученного ранее с помощью функции
cuda.get_rng_state.
Вторым параметром можно указать идентификатор устройства
(по умолчанию - текущее устройство).
Синтаксис
torch.cuda.set_rng_state(state, [device])
Пример
Сначала получим текущее состояние генератора, затем восстановим его после выполнения случайных операций:
import torch
torch.manual_seed(0)
state = torch.cuda.get_rng_state()
t = torch.randn(2, 3, device='cuda')
print(t)
torch.cuda.set_rng_state(state)
t2 = torch.randn(2, 3, device='cuda')
print(t2)
Результат выполнения кода:
tensor([
[ 1.5410, -0.2934, -2.1788],
[ 0.5684, -1.0845, -1.3986]
], device='cuda:0')
tensor([
[ 1.5410, -0.2934, -2.1788],
[ 0.5684, -1.0845, -1.3986]
], device='cuda:0')
Как видно, после восстановления состояния тензоры
t и t2 получились одинаковыми.
Пример
Восстановление состояния для конкретного устройства (например, GPU с индексом 1):
import torch
device_id = 0
torch.manual_seed(42)
state = torch.cuda.get_rng_state(device_id)
t = torch.rand(3, device='cuda:0')
print(t)
torch.cuda.set_rng_state(state, device_id)
t2 = torch.rand(3, device='cuda:0')
print(t2)
Результат выполнения кода:
tensor([0.8080, 0.3551, 0.6162], device='cuda:0')
tensor([0.8080, 0.3551, 0.6162], device='cuda:0')
Состояние было восстановлено на устройстве cuda:0,
поэтому оба тензора совпадают.
Пример
Восстановление состояния генератора после изменения потока случайных чисел:
import torch
torch.manual_seed(123)
state = torch.cuda.get_rng_state()
t1 = torch.randn(2, 2, device='cuda')
t2 = torch.randn(2, 2, device='cuda')
print("t1:", t1)
print("t2:", t2)
torch.cuda.set_rng_state(state)
t3 = torch.randn(2, 2, device='cuda')
t4 = torch.randn(2, 2, device='cuda')
print("t3:", t3)
print("t4:", t4)
Результат выполнения кода:
t1: tensor([
[ 0.5207, -0.4817],
[ 0.1106, -0.2868]
], device='cuda:0')
t2: tensor([
[ 1.1911, -0.1305],
[-1.6027, 0.3517]
], device='cuda:0')
t3: tensor([
[ 0.5207, -0.4817],
[ 0.1106, -0.2868]
], device='cuda:0')
t4: tensor([
[ 1.1911, -0.1305],
[-1.6027, 0.3517]
], device='cuda:0')
После восстановления состояния последовательность случайных чисел повторяется полностью.
Смотрите также
-
функцию
cuda.get_rng_state,
которая возвращает текущее состояние ГСЧ на GPU -
функцию
cuda.manual_seed,
которая задаёт начальное зерно для генератора на конкретном GPU -
функцию
get_rng_state,
которая работает с генератором на CPU -
функцию
manual_seed,
которая задаёт зерно для всех генераторов (CPU и GPU)