|
21 | 21 | from torch import nn |
22 | 22 |
|
23 | 23 |
|
| 24 | +def weights_and_biases_init(m: nn.Module) -> None: |
| 25 | + """ |
| 26 | + Initialize the weights and biases of the model. |
| 27 | +
|
| 28 | + Parameters |
| 29 | + ---------- |
| 30 | + m : nn.Module |
| 31 | + Module (layer) to initialize |
| 32 | +
|
| 33 | + Notes |
| 34 | + ----- |
| 35 | + This function is used to initialize the weights of the model for both convolutional |
| 36 | + and linear layers. Weights are initialized using uniform Xavier initialization |
| 37 | + while biases are set to zero. |
| 38 | +
|
| 39 | + https://github.com/gnina/libmolgrid/blob/e6d5f36f1ae03f643ca69cdec1625ac52e653f88/test/test_torch_cnn.py#L45-L48 |
| 40 | + """ |
| 41 | + if isinstance(m, nn.Conv3d) or isinstance(m, nn.Linear): |
| 42 | + nn.init.xavier_uniform_(m.weight.data) |
| 43 | + nn.init.constant_(m.bias.data, 0.0) |
| 44 | + |
| 45 | + |
24 | 46 | class Default2017(nn.Module): |
25 | 47 | """ |
26 | 48 | GNINA default2017 model architecture. |
@@ -134,13 +156,6 @@ def __init__(self, input_dims: Tuple): |
134 | 156 | ) |
135 | 157 | ) |
136 | 158 |
|
137 | | - # Xavier initialization for convolutional and linear layers |
138 | | - for m in self.modules(): |
139 | | - if isinstance(m, nn.Conv3d) or isinstance(m, nn.Linear): |
140 | | - nn.init.xavier_uniform_(m.weight.data) |
141 | | - # TODO: Initialize bias to zero? |
142 | | - # TODO: See https://github.com/gnina/libmolgrid/blob/e6d5f36f1ae03f643ca69cdec1625ac52e653f88/test/test_torch_cnn.py#L48 |
143 | | - |
144 | 159 | def forward(self, x: torch.Tensor): |
145 | 160 | """ |
146 | 161 | Parameters |
@@ -197,13 +212,6 @@ def __init__(self, input_dims: Tuple): |
197 | 212 | ) |
198 | 213 | ) |
199 | 214 |
|
200 | | - # Xavier initialization for convolutional and linear layers |
201 | | - for m in self.modules(): |
202 | | - if isinstance(m, nn.Conv3d) or isinstance(m, nn.Linear): |
203 | | - nn.init.xavier_uniform_(m.weight.data) |
204 | | - # TODO: Initialize bias to zero? |
205 | | - # TODO: See https://github.com/gnina/libmolgrid/blob/e6d5f36f1ae03f643ca69cdec1625ac52e653f88/test/test_torch_cnn.py#L48 |
206 | | - |
207 | 215 | def forward(self, x: torch.Tensor): |
208 | 216 | """ |
209 | 217 | Parameters |
@@ -368,13 +376,6 @@ def __init__(self, input_dims: Tuple): |
368 | 376 | ) |
369 | 377 | ) |
370 | 378 |
|
371 | | - # Xavier initialization for convolutional and linear layers |
372 | | - for m in self.modules(): |
373 | | - if isinstance(m, nn.Conv3d) or isinstance(m, nn.Linear): |
374 | | - nn.init.xavier_uniform_(m.weight.data) |
375 | | - # TODO: Initialize bias to zero? |
376 | | - # TODO: See https://github.com/gnina/libmolgrid/blob/e6d5f36f1ae03f643ca69cdec1625ac52e653f88/test/test_torch_cnn.py#L48 |
377 | | - |
378 | 379 | def forward(self, x: torch.Tensor): |
379 | 380 | """ |
380 | 381 | Parameters |
@@ -430,13 +431,6 @@ def __init__(self, input_dims: Tuple): |
430 | 431 | ) |
431 | 432 | ) |
432 | 433 |
|
433 | | - # Xavier initialization for convolutional and linear layers |
434 | | - for m in self.modules(): |
435 | | - if isinstance(m, nn.Conv3d) or isinstance(m, nn.Linear): |
436 | | - nn.init.xavier_uniform_(m.weight.data) |
437 | | - # TODO: Initialize bias to zero? |
438 | | - # TODO: See https://github.com/gnina/libmolgrid/blob/e6d5f36f1ae03f643ca69cdec1625ac52e653f88/test/test_torch_cnn.py#L48 |
439 | | - |
440 | 434 | def forward(self, x: torch.Tensor): |
441 | 435 | """ |
442 | 436 | Parameters |
@@ -656,11 +650,6 @@ def __init__( |
656 | 650 |
|
657 | 651 | self.features = nn.Sequential(features) |
658 | 652 |
|
659 | | - # Xavier initialization for convolutional and linear layers |
660 | | - for m in self.modules(): |
661 | | - if isinstance(m, nn.Conv3d) or isinstance(m, nn.Linear): |
662 | | - nn.init.xavier_uniform_(m.weight.data) |
663 | | - |
664 | 653 | def forward(self, x): |
665 | 654 | """ |
666 | 655 | Parameters |
@@ -719,11 +708,6 @@ def __init__( |
719 | 708 | ) |
720 | 709 | ) |
721 | 710 |
|
722 | | - # Xavier initialization for convolutional and linear layers |
723 | | - for m in self.modules(): |
724 | | - if isinstance(m, nn.Conv3d) or isinstance(m, nn.Linear): |
725 | | - nn.init.xavier_uniform_(m.weight.data) |
726 | | - |
727 | 711 | def forward(self, x): |
728 | 712 | """ |
729 | 713 | Parameters |
@@ -795,11 +779,6 @@ def __init__( |
795 | 779 | ) |
796 | 780 | ) |
797 | 781 |
|
798 | | - # Xavier initialization for convolutional and linear layers |
799 | | - for m in self.modules(): |
800 | | - if isinstance(m, nn.Conv3d) or isinstance(m, nn.Linear): |
801 | | - nn.init.xavier_uniform_(m.weight.data) |
802 | | - |
803 | 782 | def forward(self, x): |
804 | 783 | """ |
805 | 784 | Parameters |
@@ -927,11 +906,6 @@ def __init__(self, input_dims: Tuple): |
927 | 906 | ) |
928 | 907 | ) |
929 | 908 |
|
930 | | - # Xavier initialization for convolutional and linear layers |
931 | | - for m in self.modules(): |
932 | | - if isinstance(m, nn.Conv3d) or isinstance(m, nn.Linear): |
933 | | - nn.init.xavier_uniform_(m.weight.data) |
934 | | - |
935 | 909 | def forward(self, x: torch.Tensor): |
936 | 910 | """ |
937 | 911 | Parameters |
@@ -1066,11 +1040,6 @@ def __init__(self, input_dims: Tuple): |
1066 | 1040 | ) |
1067 | 1041 | ) |
1068 | 1042 |
|
1069 | | - # Xavier initialization for convolutional and linear layers |
1070 | | - for m in self.modules(): |
1071 | | - if isinstance(m, nn.Conv3d) or isinstance(m, nn.Linear): |
1072 | | - nn.init.xavier_uniform_(m.weight.data) |
1073 | | - |
1074 | 1043 | def forward(self, x: torch.Tensor): |
1075 | 1044 | """ |
1076 | 1045 | Parameters |
|
0 commit comments