Функция cuda.get_rng_state
Функция cuda.get_rng_state возвращает объект
с состоянием генератора случайных чисел
для указанного устройства CUDA.
Состояние представляет собой байтовый тензор
с информацией о текущем положении генератора.
Это полезно для воспроизводимости экспериментов
и сохранения промежуточных состояний
генератора случайных чисел.
Функция принимает один необязательный параметр
device - идентификатор GPU,
для которого нужно получить состояние.
Если параметр не указан, используется
текущее устройство.
Синтаксис
torch.cuda.get_rng_state([device])
Параметры
device (необязательный) - целое число
или объект torch.device, указывающий
на конкретный GPU. Если не указан,
используется текущее устройство.
Возвращаемое значение - байтовый тензор с состоянием генератора случайных чисел для указанного устройства.
Пример получения состояния
Давайте получим текущее состояние генератора случайных чисел для устройства по умолчанию:
import torch
state = torch.cuda.get_rng_state()
print(state)
Результат выполнения кода (байтовый тензор с состоянием):
tensor([ 0, 0, 0, ..., 0, 0, 0], dtype=torch.uint8)
Пример сохранения и восстановления состояния
Сохраним состояние генератора, сгенерируем несколько случайных чисел, затем восстановим сохранённое состояние и убедимся, что последовательность повторяется:
import torch
torch.manual_seed(42)
# Сохраняем состояние
state = torch.cuda.get_rng_state()
# Генерируем первые случайные числа
t1 = torch.rand(3)
print("First sequence:", t1)
# Восстанавливаем состояние
torch.cuda.set_rng_state(state)
# Генерируем повторно те же числа
t2 = torch.rand(3)
print("Restored sequence:", t2)
Результат выполнения кода:
First sequence: tensor([0.8823, 0.9150, 0.3829], device='cuda:0')
Restored sequence: tensor([0.8823, 0.9150, 0.3829], device='cuda:0')
Пример получения состояния для конкретного устройства
Если у вас несколько GPU, можно получить состояние для конкретного устройства:
import torch
state = torch.cuda.get_rng_state(device=0)
print(state)
Результат выполнения кода:
tensor([ 0, 0, 0, ..., 0, 0, 0], dtype=torch.uint8)
Смотрите также
-
функцию
cuda.set_rng_state,
которая восстанавливает состояние генератора для указанного устройства -
функцию
cuda.get_rng_state_all,
которая получает состояние для всех устройств одновременно -
функцию
manual_seed,
которая задаёт начальное зерно для генератора на CPU -
функцию
cuda.manual_seed,
которая задаёт начальное зерно для генератора на GPU