(m)
| 490 | |
| 491 | |
| 492 | def init_weights(m): |
| 493 | classname = m.__class__.__name__ |
| 494 | if classname.find('Conv2d') != -1 or classname.find('ConvTranspose2d') != -1: |
| 495 | nn.init.kaiming_uniform_(m.weight) |
| 496 | nn.init.zeros_(m.bias) |
| 497 | elif classname.find('BatchNorm') != -1: |
| 498 | nn.init.normal_(m.weight, 1.0, 0.02) |
| 499 | nn.init.zeros_(m.bias) |
| 500 | elif classname.find('Linear') != -1: |
| 501 | nn.init.xavier_normal_(m.weight) |
| 502 | nn.init.zeros_(m.bias) |
| 503 | |
| 504 | |
| 505 | def kl_div_with_logit(q_logit, p_logit): |
nothing calls this directly
no outgoing calls
no test coverage detected