Загрузка данных


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("Обучение завершено!")