MolecularDiffusion.modules.models.geoldm.networks¶
EGNN wrapper networks, ported from GeoLDM’s egnn/models.py (commit 03ae2031c712a1a6c1678e747bdcdc7a7560e00b).
remove_mean / remove_mean_with_mask are reused from MolecularDiffusion.utils (same implementations, no need to re-port GeoLDM’s own equivariant_diffusion/utils.py copies) per the integration plan.
Classes¶
Module Contents¶
- class MolecularDiffusion.modules.models.geoldm.networks.EGNN_decoder_QM9(in_node_nf, context_node_nf, out_node_nf, n_dims, hidden_nf=64, device='cpu', act_fn=torch.nn.SiLU(), n_layers=4, attention=False, tanh=False, mode='egnn_dynamics', norm_constant=0, inv_sublayers=2, sin_embedding=False, normalization_factor=100, aggregation_method='sum', include_charges=True)¶
Bases:
torch.nn.Module- abstractmethod forward(t, xh, node_mask, edge_mask, context=None)¶
- get_adj_matrix(n_nodes, batch_size, device)¶
- unwrap_forward()¶
- wrap_forward(node_mask, edge_mask, context)¶
- context_node_nf¶
- device = 'cpu'¶
- include_charges = 1¶
- mode = 'egnn_dynamics'¶
- n_dims¶
- num_classes¶
- class MolecularDiffusion.modules.models.geoldm.networks.EGNN_dynamics_QM9(in_node_nf, context_node_nf, n_dims, hidden_nf=64, device='cpu', act_fn=torch.nn.SiLU(), n_layers=4, attention=False, condition_time=True, tanh=False, mode='egnn_dynamics', norm_constant=0, inv_sublayers=2, sin_embedding=False, normalization_factor=100, aggregation_method='sum')¶
Bases:
torch.nn.Module- abstractmethod forward(t, xh, node_mask, edge_mask, context=None)¶
- get_adj_matrix(n_nodes, batch_size, device)¶
- unwrap_forward()¶
- wrap_forward(node_mask, edge_mask, context)¶
- condition_time = True¶
- context_node_nf¶
- device = 'cpu'¶
- mode = 'egnn_dynamics'¶
- n_dims¶
- class MolecularDiffusion.modules.models.geoldm.networks.EGNN_encoder_QM9(in_node_nf, context_node_nf, out_node_nf, n_dims, hidden_nf=64, device='cpu', act_fn=torch.nn.SiLU(), n_layers=4, attention=False, tanh=False, mode='egnn_dynamics', norm_constant=0, inv_sublayers=2, sin_embedding=False, normalization_factor=100, aggregation_method='sum', include_charges=True)¶
Bases:
torch.nn.Module- Parameters:
in_node_nf – Number of invariant features for input nodes.
- abstractmethod forward(t, xh, node_mask, edge_mask, context=None)¶
- get_adj_matrix(n_nodes, batch_size, device)¶
- unwrap_forward()¶
- wrap_forward(node_mask, edge_mask, context)¶
- context_node_nf¶
- device = 'cpu'¶
- final_mlp¶
- include_charges = 1¶
- mode = 'egnn_dynamics'¶
- n_dims¶
- num_classes¶
- out_node_nf¶