始める無料で始める

U-Net を構築する: forward メソッド

エンコーダ層とデコーダ層は定義済みなので、ここでは U-Net の forward() メソッドを実装します。入力はすでにエンコーダに通されていますが、最後のデコーダブロックを定義する必要があります。

デコーダの目的は、特徴マップをアップサンプリングして、出力の高さと幅を U-Net の入力画像と同じにすることです。これにより、ピクセル単位のセマンティックマスクを得られます。

この演習はコースの一部です

PyTorch で学ぶ画像向け Deep Learning

コースを見る

演習の手順

  • スキップ接続を作るために torch.cat() を使い、最後のデコーダブロックを定義してください。

実践的なインタラクティブ演習

このサンプルコードを完成させて、この演習に挑戦してみましょう。

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)
コードを編集して実行