waxmorph.torch.gnn.GNS#
- class waxmorph.torch.gnn.GNS(node_feature_dim, edge_feature_dim, node_latent_dim=128, edge_latent_dim=128, hidden_dim=128, num_mp_steps=10, num_mlp_layers=2, output_dims=None, activation='relu', layer_norm=True, checkpoint_processor=False)[source]#
Bases:
ModuleEncode-process-decode graph network.
Independent residual message-passing blocks consume encoded node and edge features. Named decoders map final node latents to per-node updates; defaults are
dX: 3,dP: 3, anddc: 2.num_mlp_layerscounts linear layers in every MLP.layer_normapplies to encoders and processor MLPs; decoders project directly to outputs.checkpoint_processorrecomputes processor and decoder activations during backward to reduce saved activation memory.Examples
>>> model = GNS(1, 2, hidden_dim=4, num_mp_steps=1, output_dims={"dX": 3}) >>> out = model(torch.ones(2, 1), torch.tensor([[0, 1], [1, 0]]), torch.ones(2, 2)) >>> print(sorted(out), tuple(out["dX"].shape)) ['dX'] (2, 3)
- Parameters:
- save(path)[source]#
Serialize
{"config": ..., "state_dict": ...}withtorch.save().
- classmethod load(path, **kwargs)[source]#
Load the
config/state_dictschema written bysave().Use checkpoints from trusted sources;
weights_only=Falsecan execute serialized Python code.kwargspass totorch.load().