Začněte nyníZačněte zdarma

Tvorba U-Netu: metoda forward

Teď, když máš definované vrstvy enkodéru i dekodéru, můžeš implementovat metodu forward() U-Netu. Vstupy již byly zpracovány enkodérem. Zbývá ti ale definovat poslední blok dekodéru.

Cílem dekodéru je provést upsampling příznakových map tak, aby výstup měl stejnou výšku a šířku jako vstupní obrázek U-Netu. Díky tomu získáš sémantické masky na úrovni jednotlivých pixelů.

Toto cvičení je součástí kurzu

Deep Learning pro obrázky s PyTorchem

Zobrazit kurz

Pokyny k cvičení

  • Definuj poslední blok dekodéru a pomocí torch.cat() vytvoř skip connection.

Interaktivní cvičení na vyzkoušení si v praxi

Vyzkoušejte si toto cvičení dokončením tohoto ukázkového kódu.

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)
Upravit a spustit kód