MolecularDiffusion.modules.models.geoldm.diffusion

E(n) hierarchical VAE + latent diffusion, ported from GeoLDM’s equivariant_diffusion/en_diffusion.py (commit 03ae2031c712a1a6c1678e747bdcdc7a7560e00b), lines 34-1243.

DistributionNodes, gaussian_KL, gaussian_KL_for_dimension, PredefinedNoiseSchedule, GammaNetwork are reused from MolecularDiffusion.modules.models.en_diffusion (that fork already carries the exact same implementations) instead of being re-ported. Mean-zero / Gaussian sampling utilities come from MolecularDiffusion.utils.

Two compatibility shims (not present upstream) are added directly on these ported classes so modules.tasks.diffusion.GeomMolecularGenerative’s inherited __init__/density_estimation/sample_chain machinery works without raising, per the integration plan (revision 2):

  • self.extra_norm_values = () / self.ndim_extra = 0 in EnVariationalDiffusion.__init__ / EnHierarchicalVAE.__init__.

  • forward(…) on all three classes accepts and ignores reference_indices=None, reference_freeze_mode=”all” (no-op passthrough; scaffold-freezing is out of scope for this integration).

Classes

EnHierarchicalVAE

The E(n) Hierarchical VAE Module.

EnLatentDiffusion

The E(n) Latent Diffusion Module.

EnVariationalDiffusion

The E(n) Diffusion Module.

Module Contents

class MolecularDiffusion.modules.models.geoldm.diffusion.EnHierarchicalVAE(encoder: MolecularDiffusion.modules.models.geoldm.networks.EGNN_encoder_QM9, decoder: MolecularDiffusion.modules.models.geoldm.networks.EGNN_decoder_QM9, in_node_nf: int, n_dims: int, latent_node_nf: int, kl_weight: float, norm_values=(1.0, 1.0, 1.0), norm_biases=(None, 0.0, 0.0), include_charges=True)

Bases: torch.nn.Module

The E(n) Hierarchical VAE Module.

compute_loss(x, h, node_mask, edge_mask, context)
compute_reconstruction_error(xh_rec, xh)
decode(z_xh, node_mask=None, edge_mask=None, context=None)

Computes p(x|z).

encode(x, h, node_mask=None, edge_mask=None, context=None)

Computes q(z|x).

forward(x, h, node_mask=None, edge_mask=None, context=None, reference_indices=None, reference_freeze_mode='all')

Computes the ELBO. reference_indices/reference_freeze_mode are accepted and ignored – see module docstring.

log_info()
reconstruct(x, h, node_mask=None, edge_mask=None, context=None)
sample_combined_position_feature_noise(n_samples, n_nodes, node_mask)
sample_normal(mu, sigma, node_mask, fix_noise=False)
subspace_dimensionality(node_mask)
decoder
encoder
extra_norm_values = ()
in_node_nf
include_charges = True
kl_weight
latent_node_nf
n_dims
ndim_extra = 0
norm_biases = (None, 0.0, 0.0)
norm_values = (1.0, 1.0, 1.0)
num_classes
class MolecularDiffusion.modules.models.geoldm.diffusion.EnLatentDiffusion(**kwargs)

Bases: EnVariationalDiffusion

The E(n) Latent Diffusion Module.

forward(x, h, node_mask=None, edge_mask=None, context=None, reference_indices=None, reference_freeze_mode='all')

reference_indices/reference_freeze_mode are accepted and ignored – see module docstring.

instantiate_first_stage(vae: EnHierarchicalVAE)
log_constants_p_h_given_z0(h, node_mask)
log_pxh_given_z0_without_constants(x, h, z_t, gamma_0, eps, net_out, node_mask, epsilon=1e-10)
sample(n_samples, n_nodes, node_mask, edge_mask, context, fix_noise=False)
sample_chain(n_samples, n_nodes, node_mask, edge_mask, context, keep_frames=None)
sample_p_xh_given_z0(z0, node_mask, edge_mask, context, fix_noise=False)
unnormalize_z(z, node_mask)
trainable_ae
class MolecularDiffusion.modules.models.geoldm.diffusion.EnVariationalDiffusion(dynamics, in_node_nf: int, n_dims: int, timesteps: int = 1000, parametrization='eps', noise_schedule='learned', noise_precision=0.0001, loss_type='vlb', norm_values=(1.0, 1.0, 1.0), norm_biases=(None, 0.0, 0.0), include_charges=True)

Bases: torch.nn.Module

The E(n) Diffusion Module.

SNR(gamma)
alpha(gamma, target_tensor)
check_issues_norm_values(num_stdevs=8)
compute_error(net_out, gamma_t, eps)
compute_loss(x, h, node_mask, edge_mask, context, t0_always)
compute_x_pred(net_out, zt, gamma_t)
forward(x, h, node_mask=None, edge_mask=None, context=None, reference_indices=None, reference_freeze_mode='all')

Computes the loss (type l2 or NLL) if training. And if eval then always computes NLL.

reference_indices/reference_freeze_mode are accepted and ignored – see module docstring; scaffold-freezing is out of scope for this integration.

inflate_batch_array(array, target)
kl_prior(xh, node_mask)
log_constants_p_x_given_z0(x, node_mask)
log_info()
log_pxh_given_z0_without_constants(x, h, z_t, gamma_0, eps, net_out, node_mask, epsilon=1e-10)
normalize(x, h, node_mask)
phi(x, t, node_mask, edge_mask, context)
sample(n_samples, n_nodes, node_mask, edge_mask, context, fix_noise=False)
sample_chain(n_samples, n_nodes, node_mask, edge_mask, context, keep_frames=None)
sample_combined_position_feature_noise(n_samples, n_nodes, node_mask)
sample_normal(mu, sigma, node_mask, fix_noise=False)
sample_p_xh_given_z0(z0, node_mask, edge_mask, context, fix_noise=False)
sample_p_zs_given_zt(s, t, zt, node_mask, edge_mask, context, fix_noise=False)
sigma(gamma, target_tensor)
sigma_and_alpha_t_given_s(gamma_t: torch.Tensor, gamma_s: torch.Tensor, target_tensor: torch.Tensor)
subspace_dimensionality(node_mask)
unnormalize(x, h_cat, h_int, node_mask)
unnormalize_z(z, node_mask)
T = 1000
dynamics
extra_norm_values = ()
in_node_nf
include_charges = True
loss_type = 'vlb'
n_dims
ndim_extra = 0
norm_biases = (None, 0.0, 0.0)
norm_values = (1.0, 1.0, 1.0)
num_classes
parametrization = 'eps'