| 497 | return output |
| 498 | |
| 499 | def init_weights(self): |
| 500 | for m in self.modules(): |
| 501 | if isinstance(m, nn.Conv2d): |
| 502 | # nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu') |
| 503 | nn.init.normal_(m.weight, std=0.001) |
| 504 | for name, _ in m.named_parameters(): |
| 505 | if name in ['bias']: |
| 506 | nn.init.constant_(m.bias, 0) |
| 507 | elif isinstance(m, nn.BatchNorm2d): |
| 508 | nn.init.constant_(m.weight, 1) |
| 509 | nn.init.constant_(m.bias, 0) |
| 510 | elif isinstance(m, nn.ConvTranspose2d): |
| 511 | nn.init.normal_(m.weight, std=0.001) |
| 512 | for name, _ in m.named_parameters(): |
| 513 | if name in ['bias']: |
| 514 | nn.init.constant_(m.bias, 0) |
| 515 | |
| 516 | def load_weights(self, pretrained=''): |
| 517 | pretrained = osp.expandvars(pretrained) |