waxmorph.torch.gnn.GraphNetworkBlock

waxmorph.torch.gnn.GraphNetworkBlock#

class waxmorph.torch.gnn.GraphNetworkBlock(node_latent_dim=128, edge_latent_dim=128, hidden_dim=128, num_mlp_layers=2, activation='silu', layer_norm=True)[source]#

Bases: Module

Residual edge-then-node message passing.

Each directed edge derives a message from its sender, receiver, and edge latents; incoming messages are summed at receivers. num_mlp_layers counts linear layers in each MLP, and layer_norm normalizes each MLP output.

Parameters:
  • node_latent_dim (int)

  • edge_latent_dim (int)

  • hidden_dim (int)

  • num_mlp_layers (int)

  • activation (str)

  • layer_norm (bool)

forward(node_latent, edge_latent, edge_index)[source]#

Return residual updates for node and edge latents.

edge_index is directed COO [2, E]: row 0 contains senders and row 1 receivers. Reverse contact orientations remain distinct messages.

Parameters:
Return type:

tuple[Tensor, Tensor]