| 183 | :param drop_rate: rate to use in all dropout layers |
| 184 | ''' |
| 185 | def __init__(self, node_in_dim, node_h_dim, |
| 186 | edge_in_dim, edge_h_dim, readout="sum", |
| 187 | num_layers=3, drop_rate=0.1, |
| 188 | activations=(F.relu, None), vector_gate=True): |
| 189 | |
| 190 | super().__init__() |
| 191 | self.output_dim = node_h_dim[0] |
| 192 | self.rbf_dim = edge_in_dim[0] |
| 193 | |
| 194 | self.residue_embdding = nn.Linear(node_in_dim[0], node_in_dim[0], bias=False) |
| 195 | self.W_v = nn.Sequential( |
| 196 | layer.GVPLayerNorm(node_in_dim), |
| 197 | layer.GVP(node_in_dim, node_h_dim, activations=(None, None), vector_gate=vector_gate) |
| 198 | ) |
| 199 | self.W_e = nn.Sequential( |
| 200 | layer.GVPLayerNorm(edge_in_dim), |
| 201 | layer.GVP(edge_in_dim, edge_h_dim, activations=(None, None), vector_gate=vector_gate) |
| 202 | ) |
| 203 | |
| 204 | self.layers = nn.ModuleList( |
| 205 | layer.GVPConvLayer(node_h_dim, edge_h_dim, drop_rate=drop_rate, |
| 206 | activations=activations, vector_gate=vector_gate) |
| 207 | for _ in range(num_layers)) |
| 208 | |
| 209 | ns, _ = node_h_dim |
| 210 | self.W_out = nn.Sequential( |
| 211 | layer.GVPLayerNorm(node_h_dim), |
| 212 | layer.GVP(node_h_dim, (ns, 0), activations=activations, vector_gate=vector_gate) |
| 213 | ) |
| 214 | |
| 215 | if readout == "sum": |
| 216 | self.readout = layers.SumReadout() |
| 217 | elif readout == "mean": |
| 218 | self.readout = layers.MeanReadout() |
| 219 | else: |
| 220 | raise ValueError("Unknown readout `%s`" % readout) |
| 221 | |
| 222 | def forward(self, graph, input, all_loss=None, metric=None): |
| 223 | h_node = self.residue_embdding(input) |