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¶
|
Raise if |
|
|
|
MiDi's adaptive cosine schedule: one exponent |
|
|
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
variableis NaN or nonzero outsidenode_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 possiblex_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 tos.Qtb –
(B,d0,dt)cumulative tot.
- 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
nuper 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_hotthat 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
PlaceHolderwith integer class ids;Eis upper-triangular sampled then mirrored, so symmetry holds by construction.