MolecularDiffusion.modules.models.ligandiff.edm

LigandDiff EDM. Verbatim port of the target repo’s src/edm.py.

Continuous 3D Gaussian DDPM (epsilon parametrisation) over xh = cat(coords, one_hot), noised and denoised only on the ligand_diff rows; the context rows are re-pasted clean at every step.

Classes

EDM

Equivariant diffusion model with two-mask (context/ligand) inpainting.

Module Contents

class MolecularDiffusion.modules.models.ligandiff.edm.EDM(dynamics: MolecularDiffusion.modules.models.ligandiff.egnn.Dynamics, in_node_nf: int, n_dims: int, timesteps: int = 1000, noise_schedule: str = 'learned', noise_precision: float = 0.0001, loss_type: str = 'vlb', norm_values: tuple = (1.0, 1.0, 1.0), norm_biases: tuple = (None, 0.0, 0.0))

Bases: torch.nn.Module

Equivariant diffusion model with two-mask (context/ligand) inpainting.

SNR(gamma)

Signal-to-noise ratio alpha^2 / sigma^2 given gamma.

alpha(gamma)

Compute alpha given gamma.

static cdf_standard_gaussian(x)

Standard-normal CDF.

compute_x_pred(eps_t, z_t, gamma_t, batch_seg)

Most likely prediction of x.

forward(x, h, context, ligand_diff, batch_seg, batch_size, ligand_group=None)
static gaussian_kl(q_mu_minus_p_mu_squared, q_sigma, p_sigma, d)

KL distance between two isotropic normal distributions.

static inflate_batch_array(x, batch_seg)

Sum a per-atom quantity into a per-molecule vector.

kl_prior(xh, mask, batch_seg)

KL between q(z_1 | x) and the prior N(0, 1).

log_constant_of_p_x_given_z0(mask, batch_seg, batch_size)

Normalising constant of p(x | z_0).

log_p_xh_given_z0_without_constants(h, z_0, gamma_0, eps, eps_hat, mask, batch_seg, epsilon=1e-10)

Reconstruction terms for coordinates and atom types at t=0.

noised_representation(xh, ligand_diff, context, batch_seg, gamma_t)
normalize(x, h)

Scale coordinates and features to the model’s working range.

sample_chain(x, h, context, ligand_diff, batch_seg, batch_size, ligand_group, keep_frames=None, timesteps=None)
sample_combined_position_feature_noise(x, ligand_diff)

Gaussian noise on coordinates + features, ligand rows only.

sample_normal(mu_xh, ligand_diff, sigma, batch_seg)
sample_p_xh_given_z0_only_ligandDiff(z_0, context, ligand_diff, batch_size, batch_seg, ligand_group)

Sample x, h ~ p(x, h | z_0), ligand rows only.

sample_p_zs_given_zt_only_ligandDiff(s, t, z_t, context, ligand_diff, batch_seg, ligand_group)

Sample z_s ~ p(z_s | z_t), ligand rows only.

sigma(gamma)

Compute sigma given gamma.

sigma_and_alpha_t_given_s(gamma_t: torch.Tensor, gamma_s: torch.Tensor)

alpha_t|s = alpha_t / alpha_s and the matching sigma.

unnormalize(x, h)

Inverse of normalize().

unnormalize_z(z)

Unnormalize a concatenated [x | h] latent.

T = 1000
dynamics
in_node_nf
n_dims
norm_biases = (None, 0.0, 0.0)
norm_values = (1.0, 1.0, 1.0)