Huấn luyện GAN

Deep Learning cho Ảnh với PyTorch

Michal Oleszak

Machine Learning Engineer

Mục tiêu của generator

 

 

Sơ đồ quy trình GAN.

  • Mục tiêu: Tạo ảnh giả đánh lừa discriminator
  • Ý tưởng: Dùng discriminator để đánh giá hiệu suất của generator
  • Đầu ra của generator được discriminator phân loại:
    • Thật (nhãn 1) - tốt, loss nhỏ
    • Giả (nhãn 0) - kém, loss lớn
Deep Learning cho Ảnh với PyTorch

Loss của generator

def gen_loss(gen, disc, num_images, z_dim):

noise = torch.randn(num_images, z_dim)
fake = gen(noise)
disc_pred = disc(fake)
criterion = nn.BCEWithLogitsLoss()
gen_loss = criterion( disc_pred, torch.ones_like(disc_pred) ) return gen_loss
  • Định nghĩa nhiễu ngẫu nhiên
  • Tạo ảnh giả
  • Lấy dự đoán của discriminator cho ảnh giả
  • Dùng tiêu chuẩn BCE (entropy chéo nhị phân)
  • Loss của generator: BCE giữa dự đoán của discriminator và tensor toàn số 1
Deep Learning cho Ảnh với PyTorch

Mục tiêu của discriminator

 

 

Sơ đồ quy trình GAN.

  • Mục tiêu: Phân loại đúng ảnh giả và ảnh thật
  • Đầu ra của generator phải bị phân loại là giả (nhãn 0)
  • Ảnh thật phải được phân loại là thật (nhãn 1)
Deep Learning cho Ảnh với PyTorch

Loss của discriminator

def disc_loss(gen, disc, real, num_images, z_dim):

criterion = nn.BCEWithLogitsLoss()
noise = torch.randn(num_images, z_dim)
fake = gen(noise)
disc_pred_fake = disc(fake)
fake_loss = criterion( disc_pred_fake, torch.zeros_like(disc_pred_fake) )
disc_pred_real = disc(real)
real_loss = criterion( disc_pred_real, torch.ones_like(disc_pred_real) )
disc_loss = (real_loss + fake_loss) / 2 return disc_loss
  • Định nghĩa tiêu chuẩn entropy chéo nhị phân
  • Tạo nhiễu đầu vào cho generator
  • Tạo ảnh giả
  • Lấy dự đoán của discriminator cho ảnh giả
  • Tính thành phần loss cho giả
  • Lấy dự đoán của discriminator cho ảnh thật
  • Tính thành phần loss cho thật
  • Loss cuối: trung bình giữa loss thật và giả
Deep Learning cho Ảnh với PyTorch

Vòng lặp huấn luyện GAN

for epoch in range(num_epochs):
    for real in dataloader:
        cur_batch_size = len(real)


disc_opt.zero_grad()
disc_loss = disc_loss( gen, disc, real, cur_batch_size, z_dim=16)
disc_loss.backward() disc_opt.step()
gen_opt.zero_grad()
gen_loss = gen_loss( gen, disc, cur_batch_size, z_dim=16)
gen_loss.backward() gen_opt.step()
  • Lặp qua các epoch và batch dữ liệu thật, tính kích thước batch hiện tại
  • Xóa gradient của optimizer discriminator
  • Tính loss của discriminator
  • Tính gradient và cập nhật discriminator
  • Xóa gradient của optimizer generator
  • Tính loss của generator
  • Tính gradient và cập nhật generator
Deep Learning cho Ảnh với PyTorch

Ayo berlatih!

Deep Learning cho Ảnh với PyTorch

Preparing Video For Download...