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

DiscreteUniformTransition

Uniform limit distribution over every categorical modality.

MarginalUniformTransition

Limit distribution = the dataset's own class marginals.

NoiseModel

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: NoiseModel

Uniform 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: NoiseModel

Limit distribution = the dataset’s own class marginals.

This is what every released MiDi config uses (transition: marginal), and the marginals come from train_set.graph3d_stats – not from the checkpoint, which holds nn.Module weights 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 t and return the noised batch z_t.

get_Qt(t_int: torch.Tensor) MolecularDiffusion.modules.models.midi.placeholder.PlaceHolder

One-step transition matrices from t-1 to t.

get_Qt_bar(t_int: torch.Tensor) MolecularDiffusion.modules.models.midi.placeholder.PlaceHolder

Cumulative transition matrices from 0 to t.

get_alpha_bar(t_normalized: torch.Tensor | None = None, t_int: torch.Tensor | None = None, key: str | None = None) torch.Tensor

alpha_bar_t for one modality.

get_alpha_pos_ts(t_int: torch.Tensor, s_int: torch.Tensor) torch.Tensor

alpha_t / alpha_s for positions.

get_alpha_pos_ts_sq(t_int: torch.Tensor, s_int: torch.Tensor) torch.Tensor

(alpha_t / alpha_s)^2 for positions.

get_beta(t_normalized: torch.Tensor | None = None, t_int: torch.Tensor | None = None, key: str | None = None) torch.Tensor

beta_t for 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 = T prior.

get_sigma2_bar(t_normalized: torch.Tensor | None = None, t_int: torch.Tensor | None = None, key: str | None = None) torch.Tensor

sigma_bar_t^2 for 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_t for one modality.

get_sigma_pos_sq_ratio(s_int: torch.Tensor, t_int: torch.Tensor) torch.Tensor

sigma_s^2 / sigma_t^2 for 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_T from 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')