Загрузка данных
import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader
from torchvision import datasets, transforms
from torchvision.utils import save_image
import os
# 1. Настройки (Гиперпараметры)
BATCH_SIZE = 64
LATENT_DIM = 100 # Размер "шума", из которого генерируется изображение
EPOCHS = 50
LR = 0.0002
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"Используется устройство: {DEVICE}")
# Создаем папку для сохранения результатов
os.makedirs("generated_images", exist_ok=True)
# 2. Загрузка данных (Используем MNIST - рукописные цифры)
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.5,), (0.5,)) # Нормализация в диапазон [-1, 1]
])
dataset = datasets.MNIST(root='./data', train=True, download=True, transform=transform)
dataloader = DataLoader(dataset, batch_size=BATCH_SIZE, shuffle=True)
# 3. Архитектура Генератора (Создает изображение из шума)
class Generator(nn.Module):
def __init__(self):
super(Generator, self).__init__()
self.model = nn.Sequential(
nn.Linear(LATENT_DIM, 256),
nn.LeakyReLU(0.2),
nn.Linear(256, 512),
nn.LeakyReLU(0.2),
nn.Linear(512, 1024),
nn.LeakyReLU(0.2),
nn.Linear(1024, 28 * 28), # Размер картинки MNIST 28x28
nn.Tanh() # Выход в диапазоне [-1, 1]
)
def forward(self, z):
img = self.model(z)
img = img.view(img.size(0), 1, 28, 28) # Превращаем вектор в картинку
return img
# 4. Архитектура Дискриминатора (Отличает реальное фото от фейка)
class Discriminator(nn.Module):
def __init__(self):
super(Discriminator, self).__init__()
self.model = nn.Sequential(
nn.Linear(28 * 28, 512),
nn.LeakyReLU(0.2),
nn.Linear(512, 256),
nn.LeakyReLU(0.2),
nn.Linear(256, 1),
nn.Sigmoid() # Вероятность того, что картинка настоящая (от 0 до 1)
)
def forward(self, img):
img_flat = img.view(img.size(0), -1) # Разворачиваем картинку в вектор
validity = self.model(img_flat)
return validity
# Инициализация моделей
generator = Generator().to(DEVICE)
discriminator = Discriminator().to(DEVICE)
# Функция потерь и оптимизаторы
adversarial_loss = nn.BCELoss() # Binary Cross Entropy
optimizer_G = optim.Adam(generator.parameters(), lr=LR, betas=(0.5, 0.999))
optimizer_D = optim.Adam(discriminator.parameters(), lr=LR, betas=(0.5, 0.999))
# 5. Цикл обучения
print("Начинаем обучение...")
for epoch in range(EPOCHS):
for i, (imgs, _) in enumerate(dataloader):
# --- Обучение Дискриминатора ---
real_imgs = imgs.to(DEVICE)
# Генерируем шум и создаем фейковые картинки
z = torch.randn(real_imgs.size(0), LATENT_DIM).to(DEVICE)
gen_imgs = generator(z)
optimizer_D.zero_grad()
# Потери на реальных картинках (должен выдать 1)
real_loss = adversarial_loss(discriminator(real_imgs), torch.ones(real_imgs.size(0), 1).to(DEVICE))
# Потери на фейковых картинках (должен выдать 0)
fake_loss = adversarial_loss(discriminator(gen_imgs.detach()), torch.zeros(real_imgs.size(0), 1).to(DEVICE))
d_loss = (real_loss + fake_loss) / 2
d_loss.backward()
optimizer_D.step()
# --- Обучение Генератора ---
optimizer_G.zero_grad()
# Генератор хочет, чтобы Дискриминатор подумал, что картинки настоящие (выдал 1)
g_loss = adversarial_loss(discriminator(gen_imgs), torch.ones(real_imgs.size(0), 1).to(DEVICE))
g_loss.backward()
optimizer_G.step()
# Вывод прогресса
if i % 200 == 0:
print(f"[Эпоха {epoch}/{EPOCHS}] [Батч {i}/{len(dataloader)}] "
f"[D loss: {d_loss.item():.4f}] [G loss: {g_loss.item():.4f}]")
# Сохраняем примеры сгенерированных изображений каждые 5 эпох
if epoch % 5 == 0:
save_image(gen_imgs.data[:25], f"generated_images/epoch_{epoch}.png", nrow=5, normalize=True)
print("Обучение завершено!")