Ewaluacja GAN-ów

Głębokie uczenie dla obrazów z PyTorch

Michal Oleszak

Machine Learning Engineer

Generowanie obrazów

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()
  • Utwórz tensor szumu losowego
  • Przekaż szum do generatora
  • Iteruj po liczbie obrazów
  • Wytnij i-ty obraz z tensora fake
  • Zmień kolejność wymiarów obrazu
  • Wyświetl obraz
Głębokie uczenie dla obrazów z PyTorch

Obrazy wygenerowane przez GAN

Przykładowe obrazy Pokémonów wygenerowane przez GAN.

Głębokie uczenie dla obrazów z PyTorch

Fréchet Inception Distance

  • Inception: model klasyfikacji obrazów
  • Odległość Frécheta: miara odległości między dwoma rozkładami prawdopodobieństwa
  • Fréchet Inception Distance:
    1. Użyj Inception do ekstrakcji cech z próbek obrazów rzeczywistych i fałszywych
    2. Oblicz średnie i kowariancje cech dla obu zbiorów
    3. Oblicz odległość Frécheta między rzeczywistym a fałszywym rozkładem normalnym
  • Niskie FID = fałszywe obrazy podobne do danych treningowych i zróżnicowane
  • FID < 10 = dobry wynik
Głębokie uczenie dla obrazów z PyTorch

FID w 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)
  • Zaimportuj FrechetInceptionDistance
  • Utwórz instancję metryki FID
  • Zaktualizuj metrykę fałszywymi obrazami:
    • Pomnóż przez 255
    • Rzutuj na torch.uint8
  • Analogicznie zaktualizuj metrykę rzeczywistymi obrazami
  • Oblicz wartość metryki
Głębokie uczenie dla obrazów z PyTorch

Czas na ćwiczenia!

Głębokie uczenie dla obrazów z PyTorch

Preparing Video For Download...