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

CondEquiUpdate

Update atom coordinates equivariantly, use time emb condition.

Cond_DGT_concat

Conditional Diffusion Graph Transformer with self-conditioning.

DGT_concat

Diffusion Graph Transformer with self-conditioning.

EquivariantBlock

Equivariant block based on graph relational transformer layer, without extra heads.

EquivariantMixBlock

Equivariant block based on graph relational transformer layer.

MultiCondEquiUpdate

Update atom coordinates equivariantly, use time emb condition.

Functions

modulate(x, shift, scale)

Module Contents

class MolecularDiffusion.modules.models.jodo.mol_gnn.CondEquiUpdate(hidden_dim, edge_dim, dist_dim, time_dim)

Bases: torch.nn.Module

Update 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.Module

Conditional Diffusion Graph Transformer with self-conditioning.

forward(t, xh, node_mask, edge_mask, context=None, *args, **kwargs)
Parameters:
  • t ([B] time steps in [0, 1])

  • xh ([B, N, ch1] atom feature (positions, types, formal charges))

  • node_mask ([B, N, 1])

  • edge_mask ([B*N*N, 1])

  • context

  • kwargs ('edge_x' [B, N, N, ch2])

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.Module

Diffusion Graph Transformer with self-conditioning.

forward(t, xh, node_mask, edge_mask, context=None, *args, **kwargs)
Parameters:
  • t ([B] time steps in [0, 1])

  • xh ([B, N, ch1] atom feature (positions, types, formal charges))

  • node_mask ([B, N, 1])

  • edge_mask ([B*N*N, 1])

  • context

  • kwargs ('edge_x' [B, N, N, ch2])

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.Module

Equivariant 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.Module

Equivariant 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.Module

Update 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)