Zacznij terazZacznij za darmo

Budowanie U-Netu: metoda forward

Po zdefiniowaniu warstw enkodera i dekodera możesz teraz zaimplementować metodę forward() sieci U-Net. Dane wejściowe zostały już przepuszczone przez enkoder. Twoim zadaniem jest zdefiniowanie ostatniego bloku dekodera.

Celem dekodera jest próbkowanie w górę map cech tak, aby jego wyjście miało taką samą wysokość i szerokość jak obraz wejściowy sieci U-Net. Dzięki temu uzyskasz maski semantyczne na poziomie pikseli.

To ćwiczenie jest częścią kursu

Głębokie uczenie dla obrazów z PyTorch

Zobacz kurs

Instrukcje do ćwiczenia

  • Zdefiniuj ostatni blok dekodera, używając torch.cat() do utworzenia połączenia pomijającego (skip connection).

Interaktywne ćwiczenie praktyczne

Spróbuj tego ćwiczenia, uzupełniając ten przykładowy kod.

def forward(self, x):
    x1 = self.enc1(x)
    x2 = self.enc2(self.pool(x1))
    x3 = self.enc3(self.pool(x2))
    x4 = self.enc4(self.pool(x3))

    x = self.upconv3(x4)
    x = torch.cat([x, x3], dim=1)
    x = self.dec1(x)

    x = self.upconv2(x)
    x = torch.cat([x, x2], dim=1)
    x = self.dec2(x)

    # Define the last decoder block with skip connections
    x = ____
    x = ____
    x = ____

    return self.out(x)
Edytuj i uruchom kod