MCPcopy Create free account
hub / github.com/RosettaCommons/RFdiffusion / SE3TransformerWrapper

Class SE3TransformerWrapper

rfdiffusion/SE3_network.py:12–83  ·  view source on GitHub ↗

SE(3) equivariant GCN with attention

Source from the content-addressed store, hash-verified

10from se3_transformer.model.fiber import Fiber
11
12class 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')

Callers 1

__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected