開始使用免費開始

建構 U-Net:forward 方法

在已定義編碼器與解碼器層之後,現在你可以實作 U-Net 的 forward() 方法。輸入已經替你通過編碼器。不過,你需要定義最後一個解碼器區塊。

解碼器的目標是將特徵圖上採樣,使其輸出在高度與寬度上與 U-Net 的輸入影像相同。這樣你就能得到像素層級的語意遮罩。

本練習屬於課程

使用 PyTorch 進行影像深度學習

檢視課程

練習說明

  • 定義最後一個解碼器區塊,使用 torch.cat() 形成跳接(skip connection)。

動手互動練習

試著完成這個範例程式碼,體驗一下這個練習。

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)
編輯並執行程式碼