ПочатиПочніть безкоштовно

Навчання моделі GAN

Ваша команда в PyBooks досягла гарного прогресу у створенні генератора тексту на основі Generative Adversarial Network (GAN). Ви успішно визначили мережі генератора та дискримінатора. Тепер настав час їх навчити. Завершальний крок — згенерувати трохи штучних даних і порівняти їх із реальними, щоб побачити, наскільки добре ваш GAN навчився. Ми використовували тензори як вхід, і вихід намагатиметься нагадувати вхідні тензори. Команда PyBooks зможе використати ці синтетичні дані для аналізу тексту, адже ознаки матимуть таку саму залежність, як і в текстових даних.

Генератор і дискримінатор уже ініціалізовано та збережено у змінних generator і discriminator відповідно.

У вправі ініціалізовано такі змінні:

  • seq_length = 5: довжина кожної послідовності синтетичних даних
  • num_sequences = 100: загальна кількість згенерованих послідовностей
  • num_epochs = 50: кількість повних проходів крізь набір даних
  • print_every = 10: частота виводу результатів — показувати кожні 10 епох

Ця вправа є частиною курсу

Глибоке навчання для тексту з PyTorch

Переглянути курс

Інтерактивна практична вправа

Спробуйте виконати цю вправу, доповнивши цей зразок коду.

# Define the loss function and optimizer
criterion = nn.____()
optimizer_gen = ____(generator.parameters(), lr=0.001)
optimizer_disc = ____(discriminator.parameters(), lr=0.001)

for epoch in range(num_epochs):
    for real_data in data:
      	# Unsqueezing real_data and prevent gradient recalculations
        real_data = real_data.____(0)
        noise = torch.rand((1, seq_length))
        fake_data = generator(noise)
        disc_real = discriminator(real_data)
        disc_fake = discriminator(fake_data.____())
        loss_disc = criterion(disc_real, torch.ones_like(disc_real)) + criterion(disc_fake, torch.zeros_like(disc_fake))
        optimizer_disc.zero_grad()
        loss_disc.backward()
        optimizer_disc.step()
Редагувати та запускати код