Kom igångKom igång gratis

Bygga ett U-Net: forward-metoden

Nu när encoder- och dekoderlagren är definierade kan du implementera forward()-metoden i U-Net. Indata har redan skickats genom encodern. Det återstår att definiera det sista dekoderblocket.

Dekodern ska sampla upp särdragskartorna så att utdata har samma höjd och bredd som U-Netets indatabild. Det gör det möjligt att generera semantiska masker på pixelnivå.

Den här övningen är en del av kursen

Djupinlärning för bilder med PyTorch

Visa kurs

Övningsinstruktioner

  • Definiera det sista dekoderblocket och använd torch.cat() för att skapa skip-kopplingen.

Interaktiv övning med praktiskt arbete

Testa den här övningen genom att slutföra den här exempelkoden.

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)
Redigera och kör kod