Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 6 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
71 changes: 40 additions & 31 deletions neural_lam/train_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""

Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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(
Expand Down
132 changes: 131 additions & 1 deletion tests/test_train_model_warnings.py
Original file line number Diff line number Diff line change
@@ -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(
Expand Down Expand Up @@ -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"