Метод set_state
Метод set_state класса Generator восстанавливает
состояние генератора случайных чисел PyTorch из объекта,
полученного ранее с помощью метода get_state.
Первым и единственным параметром метод принимает объект состояния
типа torch.ByteTensor. Это позволяет воспроизводить
последовательности случайных чисел, начиная с определенного момента,
что полезно для отладки и воспроизводимости экспериментов.
Синтаксис
generator.set_state(state)
Пример
Давайте создадим генератор, сохраним его состояние с помощью
метода get_state, а затем восстановим с помощью
метода set_state, чтобы получить ту же последовательность
случайных чисел:
import torch
# Создаем генератор и фиксируем начальное зерно
g = torch.Generator()
g.manual_seed(42)
# Генерируем первое случайное число
t1 = torch.rand(3, generator=g)
print("Первая последовательность:", t1)
# Сохраняем состояние генератора
state = g.get_state()
# Генерируем второе случайное число
t2 = torch.rand(3, generator=g)
print("Вторая последовательность:", t2)
# Восстанавливаем состояние и генерируем заново
g.set_state(state)
t3 = torch.rand(3, generator=g)
print("Восстановленная последовательность:", t3)
Результат выполнения кода:
Первая последовательность: tensor([0.8823, 0.9150, 0.3829])
Вторая последовательность: tensor([0.9593, 0.3904, 0.6009])
Восстановленная последовательность: tensor([0.9593, 0.3904, 0.6009])
Как видите, после восстановления состояния генератор продолжил генерировать числа с того же места, где был сохранен.
Пример
Давайте используем метод set_state для воспроизведения
эксперимента с генерацией случайных тензоров разных размеров:
import torch
torch.manual_seed(0)
g = torch.Generator()
g.manual_seed(123)
# Генерируем тензор 2x3
t1 = torch.randint(0, 10, (2, 3), generator=g)
print("Тензор 2x3:\n", t1)
# Сохраняем состояние
saved_state = g.get_state()
# Генерируем другой тензор
t2 = torch.randint(0, 10, (3, 2), generator=g)
print("Тензор 3x2:\n", t2)
# Восстанавливаем состояние
g.set_state(saved_state)
# Снова генерируем тензор 2x3
t3 = torch.randint(0, 10, (2, 3), generator=g)
print("Восстановленный тензор 2x3:\n", t3)
Результат выполнения кода:
Тензор 2x3:
tensor([[4, 9, 3],
[0, 3, 9]])
Тензор 3x2:
tensor([[7, 3],
[7, 3],
[6, 5]])
Восстановленный тензор 2x3:
tensor([[4, 9, 3],
[0, 3, 9]])
Восстановленный тензор полностью совпадает с первым, потому что мы восстановили состояние до его генерации.
Пример
Давайте создадим функцию, которая использует генератор и восстанавливает его состояние для повторяемости результатов:
import torch
def generate_random_sequence(generator, length, state=None):
# Сохраняем текущее состояние, если оно передано
if state is not None:
generator.set_state(state)
# Генерируем последовательность
return torch.randn(length, generator=generator)
# Создаем генератор
g = torch.Generator()
g.manual_seed(456)
# Сохраняем начальное состояние
initial_state = g.get_state()
# Генерируем первую последовательность
seq1 = generate_random_sequence(g, 5)
print("Последовательность 1:", seq1)
# Генерируем вторую последовательность
seq2 = generate_random_sequence(g, 5)
print("Последовательность 2:", seq2)
# Восстанавливаем начальное состояние и генерируем заново
seq3 = generate_random_sequence(g, 5, initial_state)
print("Последовательность 3 (восстановленная):", seq3)
Результат выполнения кода:
Последовательность 1: tensor([ 0.0101, -0.6902, 1.2923, 0.1531, -0.0721])
Последовательность 2: tensor([ 0.7210, -0.9806, -0.4230, -0.6689, 0.0460])
Последовательность 3 (восстановленная): tensor([ 0.0101, -0.6902, 1.2923, 0.1531, -0.0721])
Метод set_state позволяет гибко управлять состоянием
генератора, делая эксперименты полностью воспроизводимыми
даже при сложных сценариях генерации.
Смотрите также
-
метод
get_state,
который сохраняет текущее состояние генератора -
метод
manual_seed,
который устанавливает начальное зерно генератора -
метод
seed,
который инициализирует генератор случайным зерном -
класс
Generator,
который представляет генератор случайных чисел