MolecularDiffusion.modules.models.midi.diffusion_utils

The four diffusion helpers MiDi’s forward and reverse processes need.

Upstream midi/diffusion/diffusion_utils.py also carries the KL/NLL machinery (mask_distributions, posterior_distributions, gaussian_KL, SNR …). That is validation-only – it feeds the variational NLL, which is out of scope for this port – so it is not ported.

Functions

assert_correctly_masked(→ None)

Raise if variable is NaN or nonzero outside node_mask.

compute_batched_over0_posterior_distribution(...)

q(z_s | z_t, x_0) for every possible x_0.

cosine_beta_schedule_discrete(→ numpy.ndarray)

MiDi's adaptive cosine schedule: one exponent nu per modality.

onehot_float(→ torch.Tensor)

F.one_hot that returns float, since every consumer here wants it.

sample_discrete_features(...)

Multinomial-sample node types, charges and (symmetric) bonds.

Module Contents

MolecularDiffusion.modules.models.midi.diffusion_utils.assert_correctly_masked(variable: torch.Tensor, node_mask: torch.Tensor) None

Raise if variable is NaN or nonzero outside node_mask.

MolecularDiffusion.modules.models.midi.diffusion_utils.compute_batched_over0_posterior_distribution(X_t: torch.Tensor, Qt: torch.Tensor, Qsb: torch.Tensor, Qtb: torch.Tensor) torch.Tensor

q(z_s | z_t, x_0) for every possible x_0.

Parameters:
  • X_t(B,N,dt) or (B,N,N,dt).

  • Qt(B,d_{t-1},dt) one-step transition matrix.

  • Qsb(B,d0,d_{t-1}) cumulative to s.

  • Qtb(B,d0,dt) cumulative to t.

Returns:

(B, N, d0, d_{t-1}).

MolecularDiffusion.modules.models.midi.diffusion_utils.cosine_beta_schedule_discrete(timesteps: int, nu_arr: list[float], s: float = 0.008) numpy.ndarray

MiDi’s adaptive cosine schedule: one exponent nu per modality.

Returns (timesteps + 1, n_modalities) betas, ordered as ['p', 'x', 'c', 'e', 'y'].

MolecularDiffusion.modules.models.midi.diffusion_utils.onehot_float(x: torch.Tensor, num_classes: int) torch.Tensor

F.one_hot that returns float, since every consumer here wants it.

MolecularDiffusion.modules.models.midi.diffusion_utils.sample_discrete_features(probX: torch.Tensor, probE: torch.Tensor, prob_charges: torch.Tensor, node_mask: torch.Tensor) MolecularDiffusion.modules.models.midi.placeholder.PlaceHolder

Multinomial-sample node types, charges and (symmetric) bonds.

Parameters:
  • probX(B,N,dx) node-type probabilities.

  • probE(B,N,N,de) edge probabilities.

  • prob_charges(B,N,dc) charge probabilities.

  • node_mask(B,N) bool.

Returns:

A PlaceHolder with integer class ids; E is upper-triangular sampled then mirrored, so symmetry holds by construction.