Load flat graphs into pyg.HeteroData and wire into BaseGraphModel#711
Open
prajwal-tech07 wants to merge 3 commits into
Open
Load flat graphs into pyg.HeteroData and wire into BaseGraphModel#711prajwal-tech07 wants to merge 3 commits into
prajwal-tech07 wants to merge 3 commits into
Conversation
First step of the HeteroData migration (issue mllam#385), scoped to flat (non-hierarchical) graphs. - Add neural_lam.utils.heterodata with graph_dict_to_heterodata and its inverse heterodata_to_graph_dict. Grid nodes are the "data" node type and mesh nodes the "hidden" node type, following the convention used in the reference implementations linked from the issue (leifdenby/weatherduck, matschreiner/equicast); g2m/m2m/m2g become typed edges. The grid node count is taken from the datastore rather than inferred from edge indices, and the edge-feature width is read dynamically (3 for 2D, 4 for 3D). - Wire into BaseGraphModel behind a use_heterodata flag (default False): when enabled the graph is represented as a HeteroData object and the model's graph tensors are taken from it. The extracted tensors are identical to the dict ones, so the model builds and trains identically either way, and existing checkpoints remain compatible. - Thread the flag through GraphLAM. - Tests: unit tests for the conversion (structure, exact round-trip, variable edge-feature width, datastore-provided grid count, hierarchical rejection) and a GraphLAM equivalence test asserting identical parameters, graph buffers and forward output with the flag on and off.
Rename the HeteroData node types from data/hidden to grid/mesh, matching the naming used throughout the rest of neural-lam (grid_static_features, mesh_static_features, g2m/m2m/m2g). The names remain module-level constants; the reference implementations in issue mllam#385 use data/hidden, noted in the module docstring.
Extend the GraphLAM equivalence test to run a few real optimizer steps on both the dict-based and HeteroData-based models and assert identical per-step losses and identical weights afterwards, covering the issue mllam#385 requirement to carry out training with both datastructures and confirm it proceeds identically.
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
First step of the HeteroData migration (#385), scoped to flat (non-hierarchical) graphs.
neural_lam.utils.heterodatawithgraph_dict_to_heterodataand its inverseheterodata_to_graph_dict. Grid nodes are thegridnode type and mesh nodes themeshnode type, matching the terminology used throughout neural-lam (grid_static_features,mesh_static_features,g2m/m2m/m2g);g2m/m2m/m2gbecome typed edges. The grid node count is taken from the datastore rather than inferred from edge indices, and the edge-feature width is read dynamically (3 for 2D, 4 for 3D). The node/edge-type names are module-level constants; the reference implementations linked from the issue (leifdenby/weatherduck,matschreiner/equicast) usedata/hidden— happy to switch if preferred.BaseGraphModelbehind ause_heterodataflag (defaultFalse): when enabled the graph is represented as aHeteroDataobject and the model's graph tensors are taken from it. The extracted tensors are identical to the dict ones, so the model builds and trains identically either way, and existing checkpoints remain compatible.GraphLAM.Hierarchical (HiLAM) support is deliberately left as a follow-up — there is an open question on #385 about how to represent mesh levels as node/edge types.
Issue Link
Part of #385 (first implementation step; the design discussion and hierarchical follow-up remain open).
Type of change
Checklist before requesting a review
pullwith--rebaseoption if possible).Checklist for reviewers
Each PR comes with its own improvements and flaws. The reviewer should check the following:
Author checklist after completed review
reflecting type of change (add section where missing):
Checklist for assignee