| 166 | return x |
| 167 | |
| 168 | def init_params(self): |
| 169 | for m in self.modules(): |
| 170 | if isinstance(m, nn.Conv2d): |
| 171 | # init.kaiming_normal_(m.weight, mode='fan_out') |
| 172 | # init.normal(m.weight, std=0.01) |
| 173 | init.xavier_normal_(m.weight) |
| 174 | if m.bias is not None: |
| 175 | init.constant_(m.bias, 0) |
| 176 | elif isinstance(m, nn.ConvTranspose2d): |
| 177 | # init.kaiming_normal_(m.weight, mode='fan_out') |
| 178 | # init.normal_(m.weight, std=0.01) |
| 179 | init.xavier_normal_(m.weight) |
| 180 | if m.bias is not None: |
| 181 | init.constant_(m.bias, 0) |
| 182 | elif isinstance(m, nn.BatchNorm2d): # NN.BatchNorm2d |
| 183 | init.constant_(m.weight, 1) |
| 184 | init.constant_(m.bias, 0) |
| 185 | elif isinstance(m, nn.Linear): |
| 186 | init.normal_(m.weight, std=0.01) |
| 187 | if m.bias is not None: |
| 188 | init.constant_(m.bias, 0) |
| 189 | |
| 190 | |
| 191 | class FFM(nn.Module): |