MolecularDiffusion.modules.models.ditmc.build

Factories for the three published DiTMC variants.

Port of dit_mc/model_zoo/GeomDiT.py. All three shipped configs are covered:

Because embed_shortest_hops_bool is True in globals, relative_embedding_bool = relative_positional_embedding_bool or embed_shortest_hops_bool is True for ``dit_ape`` too – so both non-equivariant variants instantiate DiTEdgeEmbed and run SelfAttention with both relative-positional-encoding flags on. They differ only in embed_distances_bool.

Attributes

Functions

build_variant(...)

Build one of ape / rpe / so3 with the shipped defaults.

make_molecular_dit(...)

dit_ape / dit_rpe.

make_molecular_dit_so3(...)

dit_so3. Upstream raises unless cutoff is None.

Module Contents

MolecularDiffusion.modules.models.ditmc.build.build_variant(variant: str, node_attr_dim: int, **overrides) MolecularDiffusion.modules.models.ditmc.dit.GenerativeModel

Build one of ape / rpe / so3 with the shipped defaults.

One config serves all three variants, so overrides is filtered to the chosen factory’s signature: num_radial_basis means nothing to dit_ape, and absolute_positional_embedding_bool means nothing to dit_so3. A key that matches no variant’s signature is still an error – silently swallowing a typo’d hyperparameter is how a checkpoint quietly stops matching.

MolecularDiffusion.modules.models.ditmc.build.make_molecular_dit(node_attr_dim: int, num_layers: int = 6, num_heads: int = 8, num_features_head: int = 32, mgn_num_features: int = 256, mgn_num_layers: int = 2, mgn_activation_fn: str = 'silu', num_features_mlp: int | None = None, activation_fn_mlp: str = 'gelu', activation_fn: str = 'silu', absolute_positional_embedding_bool: bool = True, relative_positional_embedding_bool: bool = False, rpe_radial_basis_bool: bool = False, rpe_num_radial_basis: int = 8, rpe_max_frequency: float = 2 * 3.141592653589793, self_conditioning_bool: bool = False, positional_encoding_bool: bool = False, embed_shortest_hops_bool: bool = False, act_dense_correct_bool: bool = False, output: str = 'drift_and_noise') MolecularDiffusion.modules.models.ditmc.dit.GenerativeModel

dit_ape / dit_rpe.

MolecularDiffusion.modules.models.ditmc.build.make_molecular_dit_so3(node_attr_dim: int, num_layers: int = 6, num_heads: int = 8, num_features_head: int = 32, cutoff: float | None = None, mgn_num_features: int = 256, mgn_num_layers: int = 2, mgn_activation_fn: str = 'silu', max_degree: int = 1, num_features_mlp: int | None = None, activation_fn_mlp: str = 'gelu', activation_fn: str = 'silu', include_pseudotensors: bool = True, radial_basis: str = 'reciprocal_bernstein', num_radial_basis: int = 64, radial_basis_kwargs: dict | None = None, cutoff_fn: str | None = None, self_conditioning_bool: bool = False, positional_encoding_bool: bool = False, embed_shortest_hops_bool: bool = False, scale_spherical_basis_with_shortest_hops_bool: bool = False, act_dense_correct_bool: bool = False, output: str = 'drift_and_noise') MolecularDiffusion.modules.models.ditmc.dit.GenerativeModel

dit_so3. Upstream raises unless cutoff is None.

MolecularDiffusion.modules.models.ditmc.build.SHIPPED_GLOBALS
MolecularDiffusion.modules.models.ditmc.build.SHIPPED_VARIANTS
MolecularDiffusion.modules.models.ditmc.build.VARIANTS = ('ape', 'rpe', 'so3')