From d839f9de798cd057685c17e4ffa868def3fec443 Mon Sep 17 00:00:00 2001 From: Severin Magel Date: Fri, 11 Sep 2026 23:21:38 -0400 Subject: [PATCH 1/3] Accept list-typed dynedge_layer_sizes in DynEdge YAML has no tuple type, so a serialized model config reloads the inner size pairs of `dynedge_layer_sizes` as lists, tripping the `isinstance(sizes, tuple)` assertion and making any saved DynEdge config with explicit layer sizes unloadable. Coerce the inner sequences to tuples so such configs round-trip, and widen the argument type to match. Co-Authored-By: Claude Opus 4.8 Claude-Session: https://claude.ai/code/session_016pqUKMMW7ZPkhc9EChzCux --- src/graphnet/models/gnn/dynedge.py | 9 ++++++--- 1 file changed, 6 insertions(+), 3 deletions(-) diff --git a/src/graphnet/models/gnn/dynedge.py b/src/graphnet/models/gnn/dynedge.py index 995449e55..39197eded 100644 --- a/src/graphnet/models/gnn/dynedge.py +++ b/src/graphnet/models/gnn/dynedge.py @@ -1,6 +1,6 @@ """Implementation of the DynEdge GNN model architecture.""" -from typing import List, Optional, Tuple, Union +from typing import List, Optional, Sequence, Union import torch from torch import Tensor, LongTensor @@ -28,7 +28,7 @@ def __init__( *, nb_neighbours: int = 8, features_subset: Optional[Union[List[int], slice]] = None, - dynedge_layer_sizes: Optional[List[Tuple[int, ...]]] = None, + dynedge_layer_sizes: Optional[List[Sequence[int]]] = None, post_processing_layer_sizes: Optional[List[int]] = None, readout_layer_sizes: Optional[List[int]] = None, global_pooling_schemes: Optional[Union[str, List[str]]] = None, @@ -102,7 +102,10 @@ def __init__( assert isinstance(dynedge_layer_sizes, list) assert len(dynedge_layer_sizes) - assert all(isinstance(sizes, tuple) for sizes in dynedge_layer_sizes) + # YAML has no tuple type, so a serialized config reloads these inner + # size pairs as lists; coerce them back so a saved DynEdge config + # round-trips. + dynedge_layer_sizes = [tuple(sizes) for sizes in dynedge_layer_sizes] assert all(len(sizes) > 0 for sizes in dynedge_layer_sizes) assert all( all(size > 0 for size in sizes) for sizes in dynedge_layer_sizes From 353e3236cf74300683dc8a363324b7b5ddf43543 Mon Sep 17 00:00:00 2001 From: Severin Magel Date: Wed, 9 Sep 2026 22:33:05 -0400 Subject: [PATCH 2/3] Test that shipped pretrained configs construct a model Parametrizes over every config under the pretrained model directory and asserts `Model.from_config` builds it, so the committed configs cannot silently rot when a constructor argument or class name changes. Co-Authored-By: Claude Opus 4.8 Claude-Session: https://claude.ai/code/session_016pqUKMMW7ZPkhc9EChzCux --- tests/models/test_pretrained.py | 37 +++++++++++++++++++++++++++++++++ 1 file changed, 37 insertions(+) create mode 100644 tests/models/test_pretrained.py diff --git a/tests/models/test_pretrained.py b/tests/models/test_pretrained.py new file mode 100644 index 000000000..f4622db96 --- /dev/null +++ b/tests/models/test_pretrained.py @@ -0,0 +1,37 @@ +"""Unit tests for the pretrained models shipped with GraphNeT.""" + +import glob +import os + +import pytest + +from graphnet.constants import PRETRAINED_MODEL_DIR +from graphnet.models import Model +from graphnet.utilities.config import ModelConfig + + +def _config_paths() -> list: + """Return every pretrained model config shipped in the repository.""" + return sorted( + glob.glob( + os.path.join(PRETRAINED_MODEL_DIR, "**", "*.yml"), recursive=True + ) + ) + + +def _config_id(path: str) -> str: + """Return a readable test id relative to the pretrained model dir.""" + return os.path.relpath(path, PRETRAINED_MODEL_DIR) + + +@pytest.mark.parametrize("config_path", _config_paths(), ids=_config_id) +def test_pretrained_config_builds(config_path: str) -> None: + """Test that every shipped pretrained config constructs a model. + + Guards the committed configs against silent rot when a constructor + argument or class name changes elsewhere in the library. + """ + config = ModelConfig.load(config_path) + assert isinstance(config, ModelConfig) + model = Model.from_config(config, trust=True) + assert isinstance(model, Model) From 87e0a8c51a3ccfcdaf530d6676d86909032d6533 Mon Sep 17 00:00:00 2001 From: Severin Magel Date: Wed, 9 Sep 2026 22:33:45 -0400 Subject: [PATCH 3/3] Test that shipped pretrained weights load strictly For every model whose state dict is committed next to its config (the QUESO models), build from config and load the weights with a strict state-dict load, verifying the committed weights and the current architecture still agree key-for-key. Co-Authored-By: Claude Opus 4.8 Claude-Session: https://claude.ai/code/session_016pqUKMMW7ZPkhc9EChzCux --- tests/models/test_pretrained.py | 28 ++++++++++++++++++++++++++++ 1 file changed, 28 insertions(+) diff --git a/tests/models/test_pretrained.py b/tests/models/test_pretrained.py index f4622db96..215109fcf 100644 --- a/tests/models/test_pretrained.py +++ b/tests/models/test_pretrained.py @@ -24,6 +24,16 @@ def _config_id(path: str) -> str: return os.path.relpath(path, PRETRAINED_MODEL_DIR) +def _config_paths_with_state_dict() -> list: + """Return pretrained configs shipped alongside a state dict.""" + return [ + path + for path in _config_paths() + if path.endswith("_config.yml") + and os.path.exists(path.replace("_config.yml", "_state_dict.pth")) + ] + + @pytest.mark.parametrize("config_path", _config_paths(), ids=_config_id) def test_pretrained_config_builds(config_path: str) -> None: """Test that every shipped pretrained config constructs a model. @@ -35,3 +45,21 @@ def test_pretrained_config_builds(config_path: str) -> None: assert isinstance(config, ModelConfig) model = Model.from_config(config, trust=True) assert isinstance(model, Model) + + +@pytest.mark.parametrize( + "config_path", _config_paths_with_state_dict(), ids=_config_id +) +def test_pretrained_state_dict_loads(config_path: str) -> None: + """Test that shipped weights load into their model without key mismatch. + + Only applies to models whose state dict is committed next to the + config; a strict load verifies the weights and the current + architecture still agree exactly. + """ + config = ModelConfig.load(config_path) + assert isinstance(config, ModelConfig) + model = Model.from_config(config, trust=True) + + state_dict_path = config_path.replace("_config.yml", "_state_dict.pth") + model.load_state_dict(state_dict_path)