MolecularDiffusion.modules.models.ditmc.dit

DiT blocks and the top-level generative model.

Port of dit_mc/backbones/neural_network_layer.{DiTLayer,SO3DiTLayer} and backbones/generative_model.GenerativeModel. IdentityMerger is a literal no-op upstream (merger.py:13-23) and is dropped; the GenerativeLayer namedtuple collapses to a plain module list.

Classes

DiTLayer

Non-equivariant DiT block with adaLN-Zero conditioning.

GenerativeModel

conditioner -> embeddings -> DiT blocks -> readout.

SO3DiTLayer

SO(3)-equivariant DiT block.

Module Contents

class MolecularDiffusion.modules.models.ditmc.dit.DiTLayer(num_features: int, num_heads: int, num_features_mlp: int, activation_fn_mlp: str = 'gelu', activation_fn: str = 'silu', relative_embedding_qk_bool: bool = True, relative_embedding_v_bool: bool = True, act_dense_correct_bool: bool = False)

Bases: torch.nn.Module

Non-equivariant DiT block with adaLN-Zero conditioning.

act_dense_correct_bool is True in globals, so the shipped models compute Dense(act_fn(c)); the zero-init makes the modulation vector zero at init either way. Both branches are ported.

forward(graph: MolecularDiffusion.modules.models.ditmc.graphs.LatentGraph, features_nodes: torch.Tensor, features_edges: torch.Tensor | None, features_cond: torch.Tensor | None, features_time: torch.Tensor) torch.Tensor
act_dense_correct_bool = False
act_fn
ada_dense
attention
cond_norm
mlp
norm1
norm2
num_features
class MolecularDiffusion.modules.models.ditmc.dit.GenerativeModel(node_embedding: torch.nn.Module, time_embedding: torch.nn.Module, layers: torch.nn.ModuleList, readout: torch.nn.Module, edge_embedding: torch.nn.Module | None = None, conditioner: torch.nn.Module | None = None, conditioning_bool: bool = False, variant: str = 'dit')

Bases: torch.nn.Module

conditioner -> embeddings -> DiT blocks -> readout.

Returns (drift, noise) when output='drift_and_noise', otherwise a single (N, 3) tensor.

forward(time_latent: torch.Tensor, graph_latent: MolecularDiffusion.modules.models.ditmc.graphs.LatentGraph, graph_cond: MolecularDiffusion.modules.models.ditmc.graphs.CondGraph | None = None)
conditioner = None
conditioning_bool = False
edge_embedding = None
layers
node_embedding
property output: str
readout
time_embedding
variant = 'dit'
class MolecularDiffusion.modules.models.ditmc.dit.SO3DiTLayer(num_features: int, num_heads: int, num_features_mlp: int, max_degree: int, include_pseudotensors: bool, in_max_degree: int, in_num_parity: int, num_radial_basis: int, basis_max_degree: int, activation_fn_mlp: str = 'gelu', activation_fn: str = 'silu', act_dense_correct_bool: bool = False)

Bases: torch.nn.Module

SO(3)-equivariant DiT block.

Unlike DiTLayer, the input shape changes between layer 0 and the rest: the node embedding is (N, 1, 1, F), and after one block the skip connection has widened it to (N, P_out, (L+1)**2, F). in_max_degree / in_num_parity therefore differ per layer and are supplied by build.py rather than inferred lazily as Flax does.

forward(graph: MolecularDiffusion.modules.models.ditmc.graphs.LatentGraph, features_nodes: torch.Tensor, features_edges: torch.Tensor, features_cond: torch.Tensor | None, features_time: torch.Tensor) torch.Tensor
act_dense_correct_bool = False
act_fn
ada_dense
att_num_parity = 2
attention
cond_norm
in_max_degree
in_num_parity
max_degree
mid_max_degree
mid_num_parity
mlp
norm1
norm2
num_features
out_max_degree
out_num_parity
parity_output = 2