(self,x, skips)
| 131 | self.SA = SpatialAttention() |
| 132 | |
| 133 | def forward(self,x, skips): |
| 134 | |
| 135 | d4 = self.Conv_1x1(x) |
| 136 | |
| 137 | # CAM4 |
| 138 | d4 = self.CA4(d4)*d4 |
| 139 | d4 = self.SA(d4)*d4 |
| 140 | d4 = self.ConvBlock4(d4) |
| 141 | |
| 142 | # upconv3 |
| 143 | d3 = self.Up3(d4) |
| 144 | |
| 145 | # AG3 |
| 146 | x3 = self.AG3(g=d3,x=skips[0]) |
| 147 | |
| 148 | # Concat 3 |
| 149 | d3 = torch.cat((x3,d3),dim=1) |
| 150 | |
| 151 | # CAM3 |
| 152 | d3 = self.CA3(d3)*d3 |
| 153 | d3 = self.SA(d3)*d3 |
| 154 | d3 = self.ConvBlock3(d3) |
| 155 | |
| 156 | # upconv2 |
| 157 | d2 = self.Up2(d3) |
| 158 | |
| 159 | # AG2 |
| 160 | x2 = self.AG2(g=d2,x=skips[1]) |
| 161 | |
| 162 | # Concat 2 |
| 163 | d2 = torch.cat((x2,d2),dim=1) |
| 164 | |
| 165 | # CAM2 |
| 166 | d2 = self.CA2(d2)*d2 |
| 167 | d2 = self.SA(d2)*d2 |
| 168 | #print(d2.shape) |
| 169 | d2 = self.ConvBlock2(d2) |
| 170 | |
| 171 | # upconv1 |
| 172 | d1 = self.Up1(d2) |
| 173 | |
| 174 | #print(skips[2]) |
| 175 | # AG1 |
| 176 | x1 = self.AG1(g=d1,x=skips[2]) |
| 177 | |
| 178 | # Concat 1 |
| 179 | d1 = torch.cat((x1,d1),dim=1) |
| 180 | |
| 181 | # CAM1 |
| 182 | d1 = self.CA1(d1)*d1 |
| 183 | d1 = self.SA(d1)*d1 |
| 184 | d1 = self.ConvBlock1(d1) |
| 185 | return d4, d3, d2, d1 |
| 186 | |
| 187 | |
| 188 | class CASCADE_Add(nn.Module): |
nothing calls this directly
no outgoing calls
no test coverage detected