Skip to content

Commit 39985f3

Browse files
authored
Merge pull request #18 from RMeli/weightsandbiases
Remove weight initialisation from model __init__
2 parents 68ea7ce + cfbc41e commit 39985f3

2 files changed

Lines changed: 24 additions & 54 deletions

File tree

gnina/models.py

Lines changed: 22 additions & 53 deletions
Original file line numberDiff line numberDiff line change
@@ -21,6 +21,28 @@
2121
from torch import nn
2222

2323

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+
2446
class Default2017(nn.Module):
2547
"""
2648
GNINA default2017 model architecture.
@@ -134,13 +156,6 @@ def __init__(self, input_dims: Tuple):
134156
)
135157
)
136158

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-
144159
def forward(self, x: torch.Tensor):
145160
"""
146161
Parameters
@@ -197,13 +212,6 @@ def __init__(self, input_dims: Tuple):
197212
)
198213
)
199214

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-
207215
def forward(self, x: torch.Tensor):
208216
"""
209217
Parameters
@@ -368,13 +376,6 @@ def __init__(self, input_dims: Tuple):
368376
)
369377
)
370378

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-
378379
def forward(self, x: torch.Tensor):
379380
"""
380381
Parameters
@@ -430,13 +431,6 @@ def __init__(self, input_dims: Tuple):
430431
)
431432
)
432433

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-
440434
def forward(self, x: torch.Tensor):
441435
"""
442436
Parameters
@@ -656,11 +650,6 @@ def __init__(
656650

657651
self.features = nn.Sequential(features)
658652

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-
664653
def forward(self, x):
665654
"""
666655
Parameters
@@ -719,11 +708,6 @@ def __init__(
719708
)
720709
)
721710

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-
727711
def forward(self, x):
728712
"""
729713
Parameters
@@ -795,11 +779,6 @@ def __init__(
795779
)
796780
)
797781

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-
803782
def forward(self, x):
804783
"""
805784
Parameters
@@ -927,11 +906,6 @@ def __init__(self, input_dims: Tuple):
927906
)
928907
)
929908

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-
935909
def forward(self, x: torch.Tensor):
936910
"""
937911
Parameters
@@ -1066,11 +1040,6 @@ def __init__(self, input_dims: Tuple):
10661040
)
10671041
)
10681042

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-
10741043
def forward(self, x: torch.Tensor):
10751044
"""
10761045
Parameters

gnina/training.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -18,7 +18,7 @@
1818
from gnina import metrics, setup, utils
1919
from gnina.dataloaders import GriddedExamplesLoader
2020
from gnina.losses import AffinityLoss
21-
from gnina.models import models_dict
21+
from gnina.models import models_dict, weights_and_biases_init
2222

2323

2424
def options(args: Optional[List[str]] = None):
@@ -587,6 +587,7 @@ def training(args):
587587
# Create model
588588
# Select model based on architecture and affinity flag (pose vs affinity)
589589
model = models_dict[(args.model, affinity)](train_loader.dims).to(device)
590+
model.apply(weights_and_biases_init)
590591

591592
# TODO: Compile model into TorchScript
592593
# Requires model refactoring to avoid branching based on affinity

0 commit comments

Comments
 (0)