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¶
The E(n) Hierarchical VAE Module. |
|
The E(n) Latent Diffusion Module. |
|
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.ModuleThe 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:
EnVariationalDiffusionThe 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.ModuleThe 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'¶