MolecularDiffusion.modules.models.ditmc.layers¶
Shared building blocks for DiTMC. Port of dit_mc/backbones/utils.py.
Parity traps carried in here from the plan, all load-bearing:
Flax
nn.LayerNormdefaults toepsilon=1e-6; PyTorch’s default is1e-5. Every LayerNorm below pins1e-6.Flax
nn.Dense’s default kernel init islecun_normal– a truncated normal, not PyTorch’s uniform default.GaussianRandomFourierFeaturesinterleavescos, sin(stack-then-reshape), it does not concatenate halves.adaLN-Zero
Denselayers stay zero-initialized.
Attributes¶
Classes¶
Equivariant MLP: e3x |
|
Layer norm that respects degree/parity structure. |
|
|
|
Plain (elementwise-activation) MLP. Port of |
Functions¶
|
|
|
|
|
|
|
Module Contents¶
- class MolecularDiffusion.modules.models.ditmc.layers.E3MLP(in_features: int, num_features: int | collections.abc.Sequence[int], max_degree: int, num_parity: int, num_layers: int = 2, activation_fn: str = 'identity', use_bias: bool = True)¶
Bases:
torch.nn.ModuleEquivariant MLP: e3x
Denselayers with a gated activation.- forward(x: torch.Tensor) torch.Tensor¶
- activation_fn¶
- layers¶
- num_layers = 2¶
- out_features¶
- class MolecularDiffusion.modules.models.ditmc.layers.EquivariantLayerNorm(num_features: int, max_degree: int, num_parity: int, *, use_scale: bool = True, use_bias: bool = True, epsilon: float = LAYERNORM_EPS)¶
Bases:
torch.nn.ModuleLayer norm that respects degree/parity structure.
The
l=0block goes through an ordinaryLayerNorm. Thel>0channels are not mean-subtracted: the per-(parity, degree)norm over the order index is computed, its variance over the feature axis is taken, and the block is multiplied byrsqrt(var + eps)(times an optional learnablescales_lm).epsilon = 1e-6.DiTMC instantiates this only with
use_scale=False, use_bias=False, so in the shipped models it carries no parameters at all.- forward(x: torch.Tensor) torch.Tensor¶
- epsilon = 1e-06¶
- has_pseudotensors¶
- has_ylms¶
- max_degree¶
- norm00¶
- num_features¶
- num_parity¶
- use_bias = True¶
- use_scale = True¶
- class MolecularDiffusion.modules.models.ditmc.layers.GaussianRandomFourierFeatures(in_features: int, features: int, sigma: float = 1.0)¶
Bases:
torch.nn.Modulegamma(x) = [cos(2π bᵀx), sin(2π bᵀx)]interleaved.bhas shape(d, features//2)and is initializednormal(sigma).- forward(x: torch.Tensor) torch.Tensor¶
- b¶
- class MolecularDiffusion.modules.models.ditmc.layers.MLP(in_features: int, num_features: int | collections.abc.Sequence[int], num_layers: int = 2, activation_fn: str = 'identity', use_bias: bool = True, output_is_zero_at_init: bool = False)¶
Bases:
torch.nn.ModulePlain (elementwise-activation) MLP. Port of
backbones/utils.MLP.Uses
get_activation_fn(thejax.nn.<name>family), not the e3x gated ones –DiTLayeruses this andSO3DiTLayerusesE3MLP, and they are different functions with the same names.- forward(x: torch.Tensor) torch.Tensor¶
- activation_fn¶
- layers¶
- num_layers = 2¶
- out_features¶
- MolecularDiffusion.modules.models.ditmc.layers.flax_dense(in_features: int, out_features: int, *, bias: bool = True, zero_init: bool = False) torch.nn.Linear¶
flax.linen.Denseequivalent: lecun-normal kernel, zero bias.
- MolecularDiffusion.modules.models.ditmc.layers.flax_layer_norm(num_features: int, *, use_scale: bool = True, use_bias: bool = True) torch.nn.LayerNorm¶
nn.LayerNormwith Flax’s defaults (eps 1e-6, last axis only).
- MolecularDiffusion.modules.models.ditmc.layers.get_max_degree_from_tensor_e3x(x: torch.Tensor) int¶
- MolecularDiffusion.modules.models.ditmc.layers.modulate_E3adaLN(x: torch.Tensor, shift: torch.Tensor, scale: torch.Tensor) torch.Tensor¶
- MolecularDiffusion.modules.models.ditmc.layers.modulate_adaLN(x: torch.Tensor, shift: torch.Tensor, scale: torch.Tensor) torch.Tensor¶
- MolecularDiffusion.modules.models.ditmc.layers.LAYERNORM_EPS = 1e-06¶