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