MolecularDiffusion.modules.layers.gfmdiff.blocks

GFMDiff Dual-Track Transformer building blocks.

Ported from GFMDiff/models/diff/layers.py (InterBlock, DistEncoder and their private helper layers) and GFMDiff/models/diff/utils.py (RBF_Emb and the geometry helpers atom_pos_to_pair_dist, local_geometry_calc, remove_mean_with_mask, create_mask).

Only import paths were adjusted for this repo’s layout; the math is unchanged from the source repo.

Classes

Functions

atom_pos_to_pair_dist(pos)

create_mask(x)

get_angle(vec1, vec2[, keepdim])

local_geometry_calc(pos, pair_mask)

remove_mean_with_mask(x, node_mask)

Module Contents

class MolecularDiffusion.modules.layers.gfmdiff.blocks.DistEncoder(emb_dim, add_time=False)

Bases: torch.nn.Module

forward(x, time=None)
linear
rbf
rbf_param
class MolecularDiffusion.modules.layers.gfmdiff.blocks.FeedForwardNetwork(emb_dim, hidden_dim, dropout=0.0)

Bases: torch.nn.Module

forward(x)
act
dropout
emb_dim
hidden_dim
lin1
lin2
class MolecularDiffusion.modules.layers.gfmdiff.blocks.InterBlock(emb_dim, hidden_dim, num_heads, dropout, node_dropout=None, pair_dropout=None, dataset_name='qm9')

Bases: torch.nn.Module

forward(atom_emb, pair_emb, pos, coord_diff, node_mask, pair_mask, batch_size, n_nodes, tri_emb=None)
data_name = 'qm9'
dropout
emb_dim
ffn_dim
head_dim
low2high_dropout
node_attn
node_attn_dropout
node_dropout
node_ffn
node_ffn_dropout
num_heads
pair_attn
pair_attn_dropout
pair_dropout
pair_ffn
pair_ffn_dropout
pair_ln
pos_update
class MolecularDiffusion.modules.layers.gfmdiff.blocks.Low2High(emb_dim)

Bases: torch.nn.Module

forward(node_emb, node_mask)
emb_dim
lin_cat
linear1
linear2
ln
class MolecularDiffusion.modules.layers.gfmdiff.blocks.NodeTransformerLayer(emb_dim, ffn_dim, num_heads, dropout, act_dropout=None, attn_dropout=None)

Bases: torch.nn.Module

attention(atom_emb, pair_emb, pair_mask, batch_size, n_nodes)
attn_update(atom_emb, pair_emb, attn_prob, node_mask, batch_size, n_nodes)
forward(atom_emb, pair_emb, node_mask, pair_mask, batch_size, n_nodes)
act_dropout = None
attn_dropout = None
dropout
emb_dim
ffn_emb_dim
head_dim
k_e_proj
k_proj
node_ln
num_heads
pair_ln
q_proj
v_add
v_e_proj
v_proj
class MolecularDiffusion.modules.layers.gfmdiff.blocks.PosUpdate(hidden_dim, act=nn.SiLU(), tanh=False, coord_range=15.0)

Bases: torch.nn.Module

forward(x_emb, pair_emb, pos, coord_diff, node_mask, pair_mask)
act
coord_mlp
coord_range = 15.0
dist_decoder
tanh = False
class MolecularDiffusion.modules.layers.gfmdiff.blocks.RBF(centers, gamma)

Bases: torch.nn.Module

forward(x)
centers
gamma
class MolecularDiffusion.modules.layers.gfmdiff.blocks.RBF_Emb(emb_dim, rbf_param, add_time=False)

Bases: torch.nn.Module

forward(feats, time=None)
center
gamma
linear
rbf
class MolecularDiffusion.modules.layers.gfmdiff.blocks.SimLow2High(emb_dim)

Bases: torch.nn.Module

forward(node_emb, pair_mask)
emb_dim
lin_cat
linear1
linear2
ln
class MolecularDiffusion.modules.layers.gfmdiff.blocks.TriTransLayer(emb_dim, ffn_dim, num_heads, dropout, act_dropout=None, attn_dropout=None)

Bases: torch.nn.Module

forward(pair_emb, tri_emb, pair_mask, batch_size, n_nodes)
act_dropout = None
attn_dropout = None
dropout
emb_dim
ffn_emb_dim
head_dim
k_a_proj
k_e_proj
ln
num_heads
out_gate
out_proj
q_proj
tri_dim
v_a_proj
v_proj
MolecularDiffusion.modules.layers.gfmdiff.blocks.atom_pos_to_pair_dist(pos)
MolecularDiffusion.modules.layers.gfmdiff.blocks.create_mask(x)
MolecularDiffusion.modules.layers.gfmdiff.blocks.get_angle(vec1, vec2, keepdim=False)
MolecularDiffusion.modules.layers.gfmdiff.blocks.local_geometry_calc(pos, pair_mask)
MolecularDiffusion.modules.layers.gfmdiff.blocks.remove_mean_with_mask(x, node_mask)