import torch.nn.functional as F from .evolution_module import * class Evolution_Network(nn.Module): def __init__(self, n_channels, n_classes, base_c=64, bilinear=True): super(Evolution_Network, self).__init__() self.n_channels = n_channels self.n_classes = n_classes self.bilinear = bilinear base_c = base_c self.inc = DoubleConv(n_channels, base_c) self.down1 = Down(base_c * 1, base_c * 2) self.down2 = Down(base_c * 2, base_c * 4) self.down3 = Down(base_c * 4, base_c * 8) factor = 2 if bilinear else 1 self.down4 = Down(base_c * 8, base_c * 16 // factor) self.up1 = Up(base_c * 16, base_c * 8 // factor, bilinear) self.up2 = Up(base_c * 8, base_c * 4 // factor, bilinear) self.up3 = Up(base_c * 4, base_c * 2 // factor, bilinear) self.up4 = Up(base_c * 2, base_c * 1, bilinear) self.outc = OutConv(base_c * 1, n_classes) self.gamma = nn.Parameter(torch.zeros(1, n_classes, 1, 1), requires_grad=True) self.up1_v = Up(base_c * 16, base_c * 8 // factor, bilinear) self.up2_v = Up(base_c * 8, base_c * 4 // factor, bilinear) self.up3_v = Up(base_c * 4, base_c * 2 // factor, bilinear) self.up4_v = Up(base_c * 2, base_c * 1, bilinear) self.outc_v = OutConv(base_c * 1, n_classes * 2) def forward(self, x): x1 = self.inc(x) x2 = self.down1(x1) x3 = self.down2(x2) x4 = self.down3(x3) x5 = self.down4(x4) x = self.up1(x5, x4) x = self.up2(x, x3) x = self.up3(x, x2) x = self.up4(x, x1) x = self.outc(x) * self.gamma v = self.up1_v(x5, x4) v = self.up2_v(v, x3) v = self.up3_v(v, x2) v = self.up4_v(v, x1) v = self.outc_v(v) return x, v