SE(3) equivariant GCN with attention
| 10 | from se3_transformer.model.fiber import Fiber |
| 11 | |
| 12 | class SE3TransformerWrapper(nn.Module): |
| 13 | """SE(3) equivariant GCN with attention""" |
| 14 | def __init__(self, num_layers=2, num_channels=32, num_degrees=3, n_heads=4, div=4, |
| 15 | l0_in_features=32, l0_out_features=32, |
| 16 | l1_in_features=3, l1_out_features=2, |
| 17 | num_edge_features=32): |
| 18 | super().__init__() |
| 19 | # Build the network |
| 20 | self.l1_in = l1_in_features |
| 21 | # |
| 22 | fiber_edge = Fiber({0: num_edge_features}) |
| 23 | if l1_out_features > 0: |
| 24 | if l1_in_features > 0: |
| 25 | fiber_in = Fiber({0: l0_in_features, 1: l1_in_features}) |
| 26 | fiber_hidden = Fiber.create(num_degrees, num_channels) |
| 27 | fiber_out = Fiber({0: l0_out_features, 1: l1_out_features}) |
| 28 | else: |
| 29 | fiber_in = Fiber({0: l0_in_features}) |
| 30 | fiber_hidden = Fiber.create(num_degrees, num_channels) |
| 31 | fiber_out = Fiber({0: l0_out_features, 1: l1_out_features}) |
| 32 | else: |
| 33 | if l1_in_features > 0: |
| 34 | fiber_in = Fiber({0: l0_in_features, 1: l1_in_features}) |
| 35 | fiber_hidden = Fiber.create(num_degrees, num_channels) |
| 36 | fiber_out = Fiber({0: l0_out_features}) |
| 37 | else: |
| 38 | fiber_in = Fiber({0: l0_in_features}) |
| 39 | fiber_hidden = Fiber.create(num_degrees, num_channels) |
| 40 | fiber_out = Fiber({0: l0_out_features}) |
| 41 | |
| 42 | self.se3 = SE3Transformer(num_layers=num_layers, |
| 43 | fiber_in=fiber_in, |
| 44 | fiber_hidden=fiber_hidden, |
| 45 | fiber_out = fiber_out, |
| 46 | num_heads=n_heads, |
| 47 | channels_div=div, |
| 48 | fiber_edge=fiber_edge, |
| 49 | use_layer_norm=True) |
| 50 | #use_layer_norm=False) |
| 51 | |
| 52 | self.reset_parameter() |
| 53 | |
| 54 | def reset_parameter(self): |
| 55 | |
| 56 | # make sure linear layer before ReLu are initialized with kaiming_normal_ |
| 57 | for n, p in self.se3.named_parameters(): |
| 58 | if "bias" in n: |
| 59 | nn.init.zeros_(p) |
| 60 | elif len(p.shape) == 1: |
| 61 | continue |
| 62 | else: |
| 63 | if "radial_func" not in n: |
| 64 | p = init_lecun_normal_param(p) |
| 65 | else: |
| 66 | if "net.6" in n: |
| 67 | nn.init.zeros_(p) |
| 68 | else: |
| 69 | nn.init.kaiming_normal_(p, nonlinearity='relu') |