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¶
Non-equivariant DiT block with adaLN-Zero conditioning. |
|
|
|
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.ModuleNon-equivariant DiT block with adaLN-Zero conditioning.
act_dense_correct_boolisTrueinglobals, so the shipped models computeDense(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.Moduleconditioner -> embeddings -> DiT blocks -> readout.Returns
(drift, noise)whenoutput='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¶
- 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.ModuleSO(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_paritytherefore differ per layer and are supplied bybuild.pyrather 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¶