Evaluarea GAN-urilor

Deep Learning pentru imagini cu PyTorch

Michal Oleszak

Machine Learning Engineer

Generarea imaginilor

num_images_to_generate = 9
noise = torch.randn(num_images_to_generate, 16)

with torch.no_grad(): fake = gen(noise)
print(f"Generated shape: {fake.shape}")
Generated shape: torch.Size([9, 3, 96, 96])
for i in range(num_images_to_generate):

image_tensor = fake[i, :, :, :]
image_permuted = image_tensor.permute(1, 2, 0)
plt.imshow(image_permuted) plt.show()
  • Creare tensor de zgomot aleatoriu
  • Transmitere zgomot către generator
  • Iterare peste numărul de imagini
  • Selectare imaginea i din fake
  • Rearanjare dimensiuni imagine
  • Afișare imagine
Deep Learning pentru imagini cu PyTorch

Generări GAN

Eșantion de imagini Pokémon generate de un GAN.

Deep Learning pentru imagini cu PyTorch

Fréchet Inception Distance

  • Inception: model de clasificare a imaginilor
  • Distanța Fréchet: măsură de distanță între două distribuții de probabilitate
  • Fréchet Inception Distance:
    1. Extragere trăsături din imagini reale și false cu Inception
    2. Calculare medii și covarianțe ale trăsăturilor
    3. Calculare distanță Fréchet între distribuțiile normale reale și false
  • FID mic = imagini false similare datelor de antrenare și diverse
  • FID < 10 = bun
Deep Learning pentru imagini cu PyTorch

FID în PyTorch

from torchmetrics.image.fid import \
FrechetInceptionDistance


fid = FrechetInceptionDistance(feature=64)
fid.update( (fake * 255).to(torch.uint8), real=False)
fid.update( (real * 255).to(torch.uint8), real=True)
fid.compute()
tensor(7.5159)
  • Import FrechetInceptionDistance
  • Instanțiere metrică FID
  • Actualizare metrică cu imagini false:
    • Înmulțire cu 255
    • Conversie la torch.uint8
  • Actualizare metrică cu imagini reale
  • Calculare valoare metrică
Deep Learning pentru imagini cu PyTorch

Să exersăm!

Deep Learning pentru imagini cu PyTorch

Preparing Video For Download...