From 4a24927628410ef4496db6b97d208c9ec6b022f8 Mon Sep 17 00:00:00 2001 From: gitcommit90 Date: Sat, 27 Jun 2026 21:25:11 +0000 Subject: [PATCH 1/3] Fix graph LAM GNN kwarg handling Co-authored-by: Hurricane --- CHANGELOG.md | 4 ++ .../models/step_predictors/graph/graph_lam.py | 1 + neural_lam/train_model.py | 6 ++ tests/test_train_model_warnings.py | 57 ++++++++++++++++++- 4 files changed, 67 insertions(+), 1 deletion(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 300b85592..aa541d401 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -43,6 +43,10 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ### Fixed +- Allow `graph_lam` training and checkpoint reloads to accept the full set of + GNN type CLI options without passing hierarchical-only options to unsupported + constructors ([#686](https://github.com/mllam/neural-lam/issues/686)). + - Fix `RuntimeError` in `HiLAMParallel` forward pass on hierarchical graphs by offsetting edge indices into the global mesh node index space ([#679](https://github.com/mllam/neural-lam/issues/679)) - Fix `IndexError` in HiLAM forward pass by offsetting grid nodes in `zero_index_g2m`/`zero_index_m2g` by the total mesh-node count across all levels ([#642](https://github.com/mllam/neural-lam/issues/642)) @Sir-Sloth-The-Lazy diff --git a/neural_lam/models/step_predictors/graph/graph_lam.py b/neural_lam/models/step_predictors/graph/graph_lam.py index fdb12a12a..a3fbaa1ec 100644 --- a/neural_lam/models/step_predictors/graph/graph_lam.py +++ b/neural_lam/models/step_predictors/graph/graph_lam.py @@ -36,6 +36,7 @@ def __init__( output_clamping_upper: dict[str, float] | None = None, g2m_gnn_type: str = "InteractionNet", m2g_gnn_type: str = "InteractionNet", + **_kwargs: object, ) -> None: """ Initialize the GraphLAM model. diff --git a/neural_lam/train_model.py b/neural_lam/train_model.py index f98065c49..07dc8e1ef 100644 --- a/neural_lam/train_model.py +++ b/neural_lam/train_model.py @@ -62,6 +62,12 @@ def load_forecaster_module_from_checkpoint(ckpt_path, config, datastore): output_std=args.output_std, output_clamping_lower=config.training.output_clamping.lower, output_clamping_upper=config.training.output_clamping.upper, + g2m_gnn_type=getattr(args, "g2m_gnn_type", "InteractionNet"), + m2g_gnn_type=getattr(args, "m2g_gnn_type", "InteractionNet"), + mesh_up_gnn_type=getattr(args, "mesh_up_gnn_type", "InteractionNet"), + mesh_down_gnn_type=getattr( + args, "mesh_down_gnn_type", "InteractionNet" + ), ) forecaster = ARForecaster(predictor, datastore) return ForecasterModule.load_from_checkpoint( diff --git a/tests/test_train_model_warnings.py b/tests/test_train_model_warnings.py index a0b5f92a9..77a69ac68 100644 --- a/tests/test_train_model_warnings.py +++ b/tests/test_train_model_warnings.py @@ -1,11 +1,12 @@ # Standard library +from types import SimpleNamespace from unittest.mock import MagicMock, patch # Third-party import pytest # First-party -from neural_lam.train_model import main +from neural_lam.train_model import load_forecaster_module_from_checkpoint, main @pytest.mark.parametrize( @@ -87,3 +88,57 @@ def capture_init(_self, **kwargs): "create_gif" in captured_kwargs ), "create_gif was not forwarded to ForecasterModule" assert captured_kwargs["create_gif"] is True + + +def test_checkpoint_loader_restores_gnn_type_kwargs(): + """Checkpoint reload must preserve custom GNN choices from saved args.""" + args = SimpleNamespace( + model="hi_lam", + graph="hierarchical", + hidden_dim=4, + hidden_layers=1, + processor_layers=1, + mesh_aggr="sum", + num_past_forcing_steps=1, + num_future_forcing_steps=1, + output_std=False, + g2m_gnn_type="PropagationNet", + m2g_gnn_type="PropagationNet", + mesh_up_gnn_type="PropagationNet", + mesh_down_gnn_type="InteractionNet", + ) + config = SimpleNamespace( + training=SimpleNamespace( + output_clamping=SimpleNamespace(lower={}, upper={}) + ) + ) + datastore = MagicMock() + captured_kwargs = {} + + class DummyPredictor: + def __init__(self, **kwargs): + captured_kwargs.update(kwargs) + + loaded_module = MagicMock() + + with ( + patch( + "neural_lam.train_model.torch.load", + return_value={"hyper_parameters": {"args": args}}, + ), + patch("neural_lam.train_model.MODELS", {"hi_lam": DummyPredictor}), + patch("neural_lam.train_model.ARForecaster"), + patch( + "neural_lam.train_model.ForecasterModule.load_from_checkpoint", + return_value=loaded_module, + ), + ): + result = load_forecaster_module_from_checkpoint( + "model.ckpt", config, datastore + ) + + assert result is loaded_module + assert captured_kwargs["g2m_gnn_type"] == "PropagationNet" + assert captured_kwargs["m2g_gnn_type"] == "PropagationNet" + assert captured_kwargs["mesh_up_gnn_type"] == "PropagationNet" + assert captured_kwargs["mesh_down_gnn_type"] == "InteractionNet" From 640b7958a8ad817df2fd7e1dd2cd86f823537833 Mon Sep 17 00:00:00 2001 From: gitcommit90 Date: Wed, 1 Jul 2026 21:07:19 +0000 Subject: [PATCH 2/3] refactor: add build_predictor helper for explicit GNN kwargs (#686) - Route training and checkpoint reload through build_predictor - Omit hierarchical mesh GNN kwargs for graph_lam - Remove GraphLAM **_kwargs swallow - Add regression test for graph_lam kwargs --- CHANGELOG.md | 2 +- .../models/step_predictors/graph/graph_lam.py | 1 - neural_lam/train_model.py | 65 +++++++++---------- tests/test_train_model_warnings.py | 41 +++++++++++- 4 files changed, 70 insertions(+), 39 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index aa541d401..12b1f4d55 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -45,7 +45,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 - Allow `graph_lam` training and checkpoint reloads to accept the full set of GNN type CLI options without passing hierarchical-only options to unsupported - constructors ([#686](https://github.com/mllam/neural-lam/issues/686)). + constructors via a shared `build_predictor` helper ([#686](https://github.com/mllam/neural-lam/issues/686)). - Fix `RuntimeError` in `HiLAMParallel` forward pass on hierarchical graphs by offsetting edge indices into the global mesh node index space ([#679](https://github.com/mllam/neural-lam/issues/679)) diff --git a/neural_lam/models/step_predictors/graph/graph_lam.py b/neural_lam/models/step_predictors/graph/graph_lam.py index a3fbaa1ec..fdb12a12a 100644 --- a/neural_lam/models/step_predictors/graph/graph_lam.py +++ b/neural_lam/models/step_predictors/graph/graph_lam.py @@ -36,7 +36,6 @@ def __init__( output_clamping_upper: dict[str, float] | None = None, g2m_gnn_type: str = "InteractionNet", m2g_gnn_type: str = "InteractionNet", - **_kwargs: object, ) -> None: """ Initialize the GraphLAM model. diff --git a/neural_lam/train_model.py b/neural_lam/train_model.py index 07dc8e1ef..cd6d0c15e 100644 --- a/neural_lam/train_model.py +++ b/neural_lam/train_model.py @@ -23,6 +23,33 @@ from .weather_dataset import WeatherDataModule +def build_predictor(predictor_class, args, config, datastore): + """Instantiate a step predictor with explicit GNN kwargs for its model family.""" + kwargs = dict( + datastore=datastore, + graph_name=args.graph, + hidden_dim=args.hidden_dim, + hidden_layers=args.hidden_layers, + processor_layers=args.processor_layers, + mesh_aggr=args.mesh_aggr, + num_past_forcing_steps=args.num_past_forcing_steps, + num_future_forcing_steps=args.num_future_forcing_steps, + output_std=args.output_std, + output_clamping_lower=config.training.output_clamping.lower, + output_clamping_upper=config.training.output_clamping.upper, + g2m_gnn_type=getattr(args, "g2m_gnn_type", "InteractionNet"), + m2g_gnn_type=getattr(args, "m2g_gnn_type", "InteractionNet"), + ) + if getattr(args, "model", None) in ("hi_lam", "hi_lam_parallel"): + kwargs["mesh_up_gnn_type"] = getattr( + args, "mesh_up_gnn_type", "InteractionNet" + ) + kwargs["mesh_down_gnn_type"] = getattr( + args, "mesh_down_gnn_type", "InteractionNet" + ) + return predictor_class(**kwargs) + + class AdaptiveHelpFormatter(ArgumentDefaultsHelpFormatter): """``--help`` formatter that scales the column width to the terminal.""" @@ -50,25 +77,7 @@ def load_forecaster_module_from_checkpoint(ckpt_path, config, datastore): ckpt = torch.load(ckpt_path, weights_only=False) args = ckpt["hyper_parameters"]["args"] predictor_class = MODELS[args.model] - predictor = predictor_class( - datastore=datastore, - graph_name=args.graph, - hidden_dim=args.hidden_dim, - hidden_layers=args.hidden_layers, - processor_layers=args.processor_layers, - mesh_aggr=args.mesh_aggr, - num_past_forcing_steps=args.num_past_forcing_steps, - num_future_forcing_steps=args.num_future_forcing_steps, - output_std=args.output_std, - output_clamping_lower=config.training.output_clamping.lower, - output_clamping_upper=config.training.output_clamping.upper, - g2m_gnn_type=getattr(args, "g2m_gnn_type", "InteractionNet"), - m2g_gnn_type=getattr(args, "m2g_gnn_type", "InteractionNet"), - mesh_up_gnn_type=getattr(args, "mesh_up_gnn_type", "InteractionNet"), - mesh_down_gnn_type=getattr( - args, "mesh_down_gnn_type", "InteractionNet" - ), - ) + predictor = build_predictor(predictor_class, args, config, datastore) forecaster = ARForecaster(predictor, datastore) return ForecasterModule.load_from_checkpoint( ckpt_path, @@ -446,23 +455,7 @@ def main(input_args=None): # Build predictor and forecaster externally, then inject into # ForecasterModule predictor_class = MODELS[args.model] - predictor = predictor_class( - datastore=datastore, - graph_name=args.graph, - hidden_dim=args.hidden_dim, - hidden_layers=args.hidden_layers, - processor_layers=args.processor_layers, - mesh_aggr=args.mesh_aggr, - num_past_forcing_steps=args.num_past_forcing_steps, - num_future_forcing_steps=args.num_future_forcing_steps, - output_std=args.output_std, - output_clamping_lower=config.training.output_clamping.lower, - output_clamping_upper=config.training.output_clamping.upper, - g2m_gnn_type=args.g2m_gnn_type, - m2g_gnn_type=args.m2g_gnn_type, - mesh_up_gnn_type=args.mesh_up_gnn_type, - mesh_down_gnn_type=args.mesh_down_gnn_type, - ) + predictor = build_predictor(predictor_class, args, config, datastore) forecaster = ARForecaster(predictor, datastore) model = ForecasterModule( diff --git a/tests/test_train_model_warnings.py b/tests/test_train_model_warnings.py index 77a69ac68..58badf028 100644 --- a/tests/test_train_model_warnings.py +++ b/tests/test_train_model_warnings.py @@ -6,7 +6,11 @@ import pytest # First-party -from neural_lam.train_model import load_forecaster_module_from_checkpoint, main +from neural_lam.train_model import ( + build_predictor, + load_forecaster_module_from_checkpoint, + main, +) @pytest.mark.parametrize( @@ -142,3 +146,38 @@ def __init__(self, **kwargs): assert captured_kwargs["m2g_gnn_type"] == "PropagationNet" assert captured_kwargs["mesh_up_gnn_type"] == "PropagationNet" assert captured_kwargs["mesh_down_gnn_type"] == "InteractionNet" + + +def test_build_predictor_omits_hierarchical_gnn_kwargs_for_graph_lam(): + """GraphLAM must not receive hierarchical-only GNN constructor kwargs.""" + args = SimpleNamespace( + model="graph_lam", + graph="multiscale", + hidden_dim=4, + hidden_layers=1, + processor_layers=1, + mesh_aggr="sum", + num_past_forcing_steps=1, + num_future_forcing_steps=1, + output_std=False, + g2m_gnn_type="PropagationNet", + m2g_gnn_type="InteractionNet", + mesh_up_gnn_type="PropagationNet", + mesh_down_gnn_type="PropagationNet", + ) + config = SimpleNamespace( + training=SimpleNamespace( + output_clamping=SimpleNamespace(lower={}, upper={}) + ) + ) + captured_kwargs = {} + + class DummyGraphLAM: + def __init__(self, **kwargs): + captured_kwargs.update(kwargs) + + build_predictor(DummyGraphLAM, args, config, MagicMock()) + + assert "mesh_up_gnn_type" not in captured_kwargs + assert "mesh_down_gnn_type" not in captured_kwargs + assert captured_kwargs["g2m_gnn_type"] == "PropagationNet" From e70558583c0ea526a02a6dbfc12a46a8a7bc8f16 Mon Sep 17 00:00:00 2001 From: gitcommit90 Date: Mon, 13 Jul 2026 15:04:50 +0000 Subject: [PATCH 3/3] refactor: gate hierarchical GNN kwargs on BaseHiGraphModel Use issubclass(predictor_class, BaseHiGraphModel) instead of a model-name tuple so future hierarchical models get mesh_up/down kwargs automatically (sadamov review on #688). Co-authored-by: Hermes Agent --- CHANGELOG.md | 4 ++- neural_lam/train_model.py | 14 +++++++++-- tests/test_train_model_warnings.py | 40 ++++++++++++++++++++++++++++-- 3 files changed, 53 insertions(+), 5 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index f493d07f3..b9753f0bc 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -47,7 +47,9 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 - Allow `graph_lam` training and checkpoint reloads to accept the full set of GNN type CLI options without passing hierarchical-only options to unsupported - constructors via a shared `build_predictor` helper ([#686](https://github.com/mllam/neural-lam/issues/686)). + constructors via a shared `build_predictor` helper that gates hierarchical + kwargs with `issubclass(..., BaseHiGraphModel)` + ([#686](https://github.com/mllam/neural-lam/issues/686)). - Fix `RuntimeError` in `HiLAMParallel` forward pass on hierarchical graphs by offsetting edge indices into the global mesh node index space ([#679](https://github.com/mllam/neural-lam/issues/679)) diff --git a/neural_lam/train_model.py b/neural_lam/train_model.py index cd6d0c15e..bbf896d5b 100644 --- a/neural_lam/train_model.py +++ b/neural_lam/train_model.py @@ -19,7 +19,12 @@ from . import utils from .config import load_config_and_datastore from .gnn_layers import GNN_TYPES -from .models import MODELS, ARForecaster, ForecasterModule +from .models import ( + MODELS, + ARForecaster, + BaseHiGraphModel, + ForecasterModule, +) from .weather_dataset import WeatherDataModule @@ -40,7 +45,12 @@ def build_predictor(predictor_class, args, config, datastore): g2m_gnn_type=getattr(args, "g2m_gnn_type", "InteractionNet"), m2g_gnn_type=getattr(args, "m2g_gnn_type", "InteractionNet"), ) - if getattr(args, "model", None) in ("hi_lam", "hi_lam_parallel"): + # Gate on class hierarchy so future hierarchical models are covered + # without maintaining a model-name list. Non-class callables (e.g. + # MagicMock in unit tests) are treated as non-hierarchical. + if isinstance(predictor_class, type) and issubclass( + predictor_class, BaseHiGraphModel + ): kwargs["mesh_up_gnn_type"] = getattr( args, "mesh_up_gnn_type", "InteractionNet" ) diff --git a/tests/test_train_model_warnings.py b/tests/test_train_model_warnings.py index 58badf028..31ebebd36 100644 --- a/tests/test_train_model_warnings.py +++ b/tests/test_train_model_warnings.py @@ -6,6 +6,7 @@ import pytest # First-party +from neural_lam.models import BaseHiGraphModel from neural_lam.train_model import ( build_predictor, load_forecaster_module_from_checkpoint, @@ -119,8 +120,9 @@ def test_checkpoint_loader_restores_gnn_type_kwargs(): datastore = MagicMock() captured_kwargs = {} - class DummyPredictor: + class DummyHiPredictor(BaseHiGraphModel): def __init__(self, **kwargs): + # Capture constructor kwargs without running full model init. captured_kwargs.update(kwargs) loaded_module = MagicMock() @@ -130,7 +132,7 @@ def __init__(self, **kwargs): "neural_lam.train_model.torch.load", return_value={"hyper_parameters": {"args": args}}, ), - patch("neural_lam.train_model.MODELS", {"hi_lam": DummyPredictor}), + patch("neural_lam.train_model.MODELS", {"hi_lam": DummyHiPredictor}), patch("neural_lam.train_model.ARForecaster"), patch( "neural_lam.train_model.ForecasterModule.load_from_checkpoint", @@ -181,3 +183,37 @@ def __init__(self, **kwargs): assert "mesh_up_gnn_type" not in captured_kwargs assert "mesh_down_gnn_type" not in captured_kwargs assert captured_kwargs["g2m_gnn_type"] == "PropagationNet" + + +def test_build_predictor_adds_hierarchical_kwargs_for_base_hi_graph_subclass(): + """Future BaseHiGraphModel subclasses get hierarchical GNN kwargs.""" + args = SimpleNamespace( + model="future_hi_model", + graph="hierarchical", + hidden_dim=4, + hidden_layers=1, + processor_layers=1, + mesh_aggr="sum", + num_past_forcing_steps=1, + num_future_forcing_steps=1, + output_std=False, + g2m_gnn_type="InteractionNet", + m2g_gnn_type="InteractionNet", + mesh_up_gnn_type="PropagationNet", + mesh_down_gnn_type="PropagationNet", + ) + config = SimpleNamespace( + training=SimpleNamespace( + output_clamping=SimpleNamespace(lower={}, upper={}) + ) + ) + captured_kwargs = {} + + class DummyFutureHiModel(BaseHiGraphModel): + def __init__(self, **kwargs): + captured_kwargs.update(kwargs) + + build_predictor(DummyFutureHiModel, args, config, MagicMock()) + + assert captured_kwargs["mesh_up_gnn_type"] == "PropagationNet" + assert captured_kwargs["mesh_down_gnn_type"] == "PropagationNet"