MolecularDiffusion.modules.models.midi.noise_model¶
MiDi’s joint noise process: D3PM on the categoricals, VP-SDE on positions.
Ported from midi/diffusion/noise_model.py with one deliberate change: the
constructor takes plain arguments instead of the upstream Hydra cfg
object, so nothing here depends on MiDi’s config tree.
This is not an nn.Module. The transition matrices and marginals are
plain tensors derived from dataset statistics, which is why they are absent
from every released MiDi checkpoint and must be rebuilt at construction time
from train_set.graph3d_stats.
Attributes¶
Classes¶
Uniform limit distribution over every categorical modality. |
|
Limit distribution = the dataset's own class marginals. |
|
Forward/reverse noise process shared by both transition variants. |
Module Contents¶
- class MolecularDiffusion.modules.models.midi.noise_model.DiscreteUniformTransition(output_dims: MolecularDiffusion.modules.models.midi.placeholder.Dims, nu: dict[str, float], diffusion_steps: int = 500, noise_schedule: str = 'cosine')¶
Bases:
NoiseModelUniform limit distribution over every categorical modality.
- E_classes¶
- E_marginals¶
- Pcharges¶
- Pe¶
- Px¶
- Py¶
- X_classes¶
- X_marginals¶
- charges_classes¶
- charges_marginals¶
- y_classes¶
- y_marginals¶
- class MolecularDiffusion.modules.models.midi.noise_model.MarginalUniformTransition(x_marginals: torch.Tensor, e_marginals: torch.Tensor, charges_marginals: torch.Tensor, y_classes: int, nu: dict[str, float], diffusion_steps: int = 500, noise_schedule: str = 'cosine')¶
Bases:
NoiseModelLimit distribution = the dataset’s own class marginals.
This is what every released MiDi config uses (
transition: marginal), and the marginals come fromtrain_set.graph3d_stats– not from the checkpoint, which holdsnn.Moduleweights only.- E_classes¶
- E_marginals¶
- Pcharges¶
- Pe¶
- Px¶
- Py¶
- X_classes¶
- X_marginals¶
- charges_classes¶
- charges_marginals¶
- y_classes¶
- y_marginals¶
- class MolecularDiffusion.modules.models.midi.noise_model.NoiseModel(nu: dict[str, float], diffusion_steps: int = 500, noise_schedule: str = 'cosine')¶
Forward/reverse noise process shared by both transition variants.
- apply_noise(dense_data: MolecularDiffusion.modules.models.midi.placeholder.PlaceHolder) MolecularDiffusion.modules.models.midi.placeholder.PlaceHolder¶
Sample
tand return the noised batchz_t.
- get_Qt(t_int: torch.Tensor) MolecularDiffusion.modules.models.midi.placeholder.PlaceHolder¶
One-step transition matrices from
t-1tot.
- get_Qt_bar(t_int: torch.Tensor) MolecularDiffusion.modules.models.midi.placeholder.PlaceHolder¶
Cumulative transition matrices from
0tot.
- get_alpha_bar(t_normalized: torch.Tensor | None = None, t_int: torch.Tensor | None = None, key: str | None = None) torch.Tensor¶
alpha_bar_tfor one modality.
- get_alpha_pos_ts(t_int: torch.Tensor, s_int: torch.Tensor) torch.Tensor¶
alpha_t / alpha_sfor positions.
- get_alpha_pos_ts_sq(t_int: torch.Tensor, s_int: torch.Tensor) torch.Tensor¶
(alpha_t / alpha_s)^2for positions.
- get_beta(t_normalized: torch.Tensor | None = None, t_int: torch.Tensor | None = None, key: str | None = None) torch.Tensor¶
beta_tfor one modality.
- get_gamma(t_normalized: torch.Tensor | None = None, t_int: torch.Tensor | None = None, key: str | None = None) torch.Tensor¶
gamma_t(log SNR) for one modality.
- get_limit_dist() MolecularDiffusion.modules.models.midi.placeholder.PlaceHolder¶
Smoothed marginals used as the
t = Tprior.
- get_sigma2_bar(t_normalized: torch.Tensor | None = None, t_int: torch.Tensor | None = None, key: str | None = None) torch.Tensor¶
sigma_bar_t^2for one modality.
- get_sigma_bar(t_normalized: torch.Tensor | None = None, t_int: torch.Tensor | None = None, key: str | None = None) torch.Tensor¶
sigma_bar_tfor one modality.
- get_sigma_pos_sq_ratio(s_int: torch.Tensor, t_int: torch.Tensor) torch.Tensor¶
sigma_s^2 / sigma_t^2for positions.
- get_x_pos_prefactor(s_int: torch.Tensor, t_int: torch.Tensor) torch.Tensor¶
a_s (s_t^2 - a_ts^2 s_s^2) / s_t^2.
- move_P_device(tensor: torch.Tensor) MolecularDiffusion.modules.models.midi.placeholder.PlaceHolder¶
Transition matrices on
tensor’s device.
- sample_limit_dist(node_mask: torch.Tensor) MolecularDiffusion.modules.models.midi.placeholder.PlaceHolder¶
Draw
z_Tfrom the limit (marginal) distribution.
- sample_zs_from_zt_and_pred(z_t: MolecularDiffusion.modules.models.midi.placeholder.PlaceHolder, pred: MolecularDiffusion.modules.models.midi.placeholder.PlaceHolder, s_int: torch.Tensor) MolecularDiffusion.modules.models.midi.placeholder.PlaceHolder¶
One reverse step:
z_s ~ p(z_s | z_t)given the denoiser output.
- E_classes = 0¶
- E_marginals: torch.Tensor | None = None¶
- Pcharges: torch.Tensor | None = None¶
- Pe: torch.Tensor | None = None¶
- Px: torch.Tensor | None = None¶
- Py: torch.Tensor | None = None¶
- T = 500¶
- X_classes = 0¶
- X_marginals: torch.Tensor | None = None¶
- charges_classes = 0¶
- charges_marginals: torch.Tensor | None = None¶
- inverse_mapping¶
- mapping = ['p', 'x', 'c', 'e', 'y']¶
- noise_schedule = 'cosine'¶
- nu_arr¶
- timesteps = 500¶
- y_classes = 0¶
- y_marginals: torch.Tensor | None = None¶
- MolecularDiffusion.modules.models.midi.noise_model.MODALITIES = ('p', 'x', 'c', 'e', 'y')¶