MolecularDiffusion.modules.models.jodo.mol_gnn¶
JODO’s Diffusion Graph Transformer (DGT), ported from the target repo.
Source: others/JODO/models/mol_gnn.py, commit 0e6326a. Ported verbatim apart from three mechanical changes:
the @utils.register_model decorators are dropped (Hydra _target_ is this platform’s registry);
from .layers import * becomes explicit imports and the two eval(name) lookups become a small dict, so ruff/mypy can see the names;
DGT_concat_2D and DGT_concat_sim are not ported – the 2D-only variant has no coordinates (nothing for the platform’s 3D pipeline) and _sim is an ablation with no released weights.
Both DGT_concat (unconditional) and Cond_DGT_concat (property-conditional) take one config namespace with .data and .model attributes, exactly as upstream, so the ported code stays diffable against JODO.
Classes¶
Update atom coordinates equivariantly, use time emb condition. |
|
Conditional Diffusion Graph Transformer with self-conditioning. |
|
Diffusion Graph Transformer with self-conditioning. |
|
Equivariant block based on graph relational transformer layer, without extra heads. |
|
Equivariant block based on graph relational transformer layer. |
|
Update atom coordinates equivariantly, use time emb condition. |
Functions¶
|
Module Contents¶
- class MolecularDiffusion.modules.models.jodo.mol_gnn.CondEquiUpdate(hidden_dim, edge_dim, dist_dim, time_dim)¶
Bases:
torch.nn.ModuleUpdate atom coordinates equivariantly, use time emb condition.
- forward(h, pos, edge_index, edge_attr, dist, time_emb)¶
- coord_mlp¶
- coord_norm¶
- input_lin¶
- ln¶
- time_mlp¶
- class MolecularDiffusion.modules.models.jodo.mol_gnn.Cond_DGT_concat(config)¶
Bases:
torch.nn.ModuleConditional Diffusion Graph Transformer with self-conditioning.
- forward(t, xh, node_mask, edge_mask, context=None, *args, **kwargs)¶
- CoM¶
- cond_lin¶
- cond_mlp¶
- dist_dim¶
- edge_emb¶
- edge_exist_mlp¶
- edge_th¶
- edge_type_mlp¶
- node_emb¶
- node_pred_mlp¶
- pred_data¶
- spatial_cut_off¶
- class MolecularDiffusion.modules.models.jodo.mol_gnn.DGT_concat(config)¶
Bases:
torch.nn.ModuleDiffusion Graph Transformer with self-conditioning.
- forward(t, xh, node_mask, edge_mask, context=None, *args, **kwargs)¶
- CoM¶
- dist_dim¶
- edge_emb¶
- edge_exist_mlp¶
- edge_th¶
- edge_type_mlp¶
- node_emb¶
- node_pred_mlp¶
- pred_data¶
- spatial_cut_off¶
- class MolecularDiffusion.modules.models.jodo.mol_gnn.EquivariantBlock(node_dim, edge_dim, time_dim, num_heads, cond_time, dist_gbf, softmax_inf, mlp_ratio=2, act=nn.SiLU(), dropout=0.0, gbf_name='GaussianLayer')¶
Bases:
torch.nn.ModuleEquivariant block based on graph relational transformer layer, without extra heads.
- forward(pos, h, edge_attr, edge_index, node_mask, node_time_emb=None, edge_time_emb=None)¶
- Params:
pos: [B*N, 3] h: [B*N, hid_dim] edge_attr: [N_edge, edge_hid_dim] edge_index: [2, N_edge] node_mask: [B*N, 1] extra_heads: [N_edge, extra_heads]
- act¶
- attn_mpnn¶
- cond_time¶
- dist_gbf¶
- dropout¶
- edge_emb¶
- edge_time_mlp¶
- equi_update¶
- ff_linear1¶
- ff_linear2¶
- ff_linear3¶
- ff_linear4¶
- node2edge_lin¶
- node_time_mlp¶
- norm1_edge¶
- norm1_node¶
- norm2_edge¶
- norm2_node¶
- class MolecularDiffusion.modules.models.jodo.mol_gnn.EquivariantMixBlock(node_dim, edge_dim, time_dim, num_extra_heads, num_heads, cond_time, dist_gbf, softmax_inf, mlp_ratio=2, act=nn.SiLU(), dropout=0.0, gbf_name='GaussianLayer', trans_name='TransMixLayer')¶
Bases:
torch.nn.ModuleEquivariant block based on graph relational transformer layer.
- forward(pos, h, edge_attr, edge_index, node_mask, extra_heads, node_time_emb=None, edge_time_emb=None)¶
- Params:
pos: [B*N, 3] h: [B*N, hid_dim] edge_attr: [N_edge, edge_hid_dim] edge_index: [2, N_edge] node_mask: [B*N, 1] extra_heads: [N_edge, extra_heads]
- act¶
- attn_mpnn¶
- cond_time¶
- dist_gbf¶
- dropout¶
- edge_emb¶
- edge_time_mlp¶
- equi_update¶
- ff_linear1¶
- ff_linear2¶
- ff_linear3¶
- ff_linear4¶
- node2edge_lin¶
- node_time_mlp¶
- norm1_edge¶
- norm1_node¶
- norm2_edge¶
- norm2_node¶
- class MolecularDiffusion.modules.models.jodo.mol_gnn.MultiCondEquiUpdate(hidden_dim, edge_dim, dist_dim, time_dim, extra_heads)¶
Bases:
torch.nn.ModuleUpdate atom coordinates equivariantly, use time emb condition.
- forward(h, pos, edge_index, edge_attr, dist, time_emb, adj_extra)¶
- coord_mlp¶
- coord_norm¶
- input_lin¶
- ln¶
- time_mlp¶
- MolecularDiffusion.modules.models.jodo.mol_gnn.modulate(x, shift, scale)¶