diff --git a/CHANGELOG.md b/CHANGELOG.md index fd123dee..b9753f0b 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -45,6 +45,12 @@ 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 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)) - 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/train_model.py b/neural_lam/train_model.py index f98065c4..bbf896d5 100644 --- a/neural_lam/train_model.py +++ b/neural_lam/train_model.py @@ -19,10 +19,47 @@ 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 +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"), + ) + # 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" + ) + 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,19 +87,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, - ) + predictor = build_predictor(predictor_class, args, config, datastore) forecaster = ARForecaster(predictor, datastore) return ForecasterModule.load_from_checkpoint( ckpt_path, @@ -440,23 +465,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 a0b5f92a..31ebebd3 100644 --- a/tests/test_train_model_warnings.py +++ b/tests/test_train_model_warnings.py @@ -1,11 +1,17 @@ # 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.models import BaseHiGraphModel +from neural_lam.train_model import ( + build_predictor, + load_forecaster_module_from_checkpoint, + main, +) @pytest.mark.parametrize( @@ -87,3 +93,127 @@ 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 DummyHiPredictor(BaseHiGraphModel): + def __init__(self, **kwargs): + # Capture constructor kwargs without running full model init. + 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": DummyHiPredictor}), + 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" + + +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" + + +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"