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¶
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.ModuleEquivariant diffusion model with two-mask (context/ligand) inpainting.
- SNR(gamma)¶
Signal-to-noise ratio
alpha^2 / sigma^2given 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 priorN(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_sand 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)¶