waxmorph.jax.gnn.GraphNetworkBlock

waxmorph.jax.gnn.GraphNetworkBlock#

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

Bases: Module

Residual edge and node updates over directed COO messages.

Each sender -> receiver row has a separate edge latent. Incoming real-edge messages are summed by receiver; num_edges masks padding. key initializes the two MLPs independently.

Parameters:
__call__(node_latent, edge_latent, edge_index, num_edges=None)[source]#

Return residual node and edge latents.

edge_index is directed COO [2, E] (senders, receivers). Entries at indices greater than or equal to num_edges contribute zero messages.

Parameters:
Return type:

tuple[Array, Array]