diff --git a/CHANGELOG.md b/CHANGELOG.md index 4ff85253..a7fdb382 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -10,6 +10,8 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ### Added - Add latent encoder/decoder modules and the `GraphEFM` (hierarchical) / `GraphEFMMultiScale` (flat) step predictors for the Graph-EFM ensemble forecasting model. [\#648](https://github.com/mllam/neural-lam/pull/648) @Sir-Sloth-The-Lazy +- Add `neural_lam.create_graph_with_wmg` CLI which builds `keisler`, `graphcast` and `hierarchical` graphs with [weather-model-graphs](https://github.com/mllam/weather-model-graphs), deprecating `neural_lam.create_graph`. [\#596](https://github.com/mllam/neural-lam/pull/596) @prajwal-tech07 + - Add `--num_sanity_val_steps` CLI argument to control sanity validation steps before training (#694) - Add `--train_steps_to_log` CLI option to log training loss for individual unroll steps, and deduplicate common prediction and loss computation steps across loops [\#674](https://github.com/mllam/neural-lam/issues/674) @GiGiKoneti diff --git a/README.md b/README.md index 35d029a6..996d51e7 100644 --- a/README.md +++ b/README.md @@ -399,6 +399,9 @@ python -m neural_lam.datastore.npyfilesmeps.compute_standardization_stats **Note:** The `create_graph` command below is deprecated and will be removed +> in a future release. Please use `create_graph_with_wmg` (see below) instead. + Run `python -m neural_lam.create_graph` with suitable options to generate the graph you want to use (see `python -m neural_lam.create_graph --help` for a list of options). The graphs used for the different models in the [paper](#graph-based-neural-weather-prediction-for-limited-area-modeling) can be created as: @@ -408,6 +411,26 @@ The graphs used for the different models in the [paper](#graph-based-neural-weat The graph-related files are stored in a directory called `graphs`. +### Graph creation with weather-model-graphs + +The recommended way to create graphs is with the `create_graph_with_wmg` +command, which delegates graph construction to +[weather-model-graphs](https://github.com/mllam/weather-model-graphs): + +```bash +python -m neural_lam.create_graph_with_wmg --config_path --archetype +``` + +Available archetypes: + +* **keisler** (default): `python -m neural_lam.create_graph_with_wmg --config_path --archetype keisler` +* **graphcast**: `python -m neural_lam.create_graph_with_wmg --config_path --archetype graphcast` +* **hierarchical**: `python -m neural_lam.create_graph_with_wmg --config_path --archetype hierarchical` + +Run `python -m neural_lam.create_graph_with_wmg --help` for the full list of +options (e.g. `--mesh_node_distance`, `--mesh_grid_distance_ratio`, +`--level_refinement_factor`, `--max_num_levels`). + ## Logging your experiments ### Weights & Biases Integration diff --git a/neural_lam/create_graph.py b/neural_lam/create_graph.py index 19734d72..8b8eae01 100644 --- a/neural_lam/create_graph.py +++ b/neural_lam/create_graph.py @@ -2,6 +2,7 @@ # Standard library import os +import warnings from argparse import ArgumentDefaultsHelpFormatter, ArgumentParser from typing import Optional @@ -910,6 +911,14 @@ def cli(input_args: Optional[list[str]] = None) -> None: Argument list forwarded to :class:`argparse.ArgumentParser`. When ``None``, ``sys.argv`` is used. """ + warnings.warn( + "create_graph.py is deprecated and will be removed in a future " + "version. Use create_graph_with_wmg.py instead, which delegates " + "graph creation to weather-model-graphs (wmg). See " + "https://github.com/mllam/neural-lam/issues/384 for details.", + DeprecationWarning, + stacklevel=2, + ) parser = ArgumentParser( description="Graph generation for neural-lam", formatter_class=ArgumentDefaultsHelpFormatter, diff --git a/neural_lam/create_graph_with_wmg.py b/neural_lam/create_graph_with_wmg.py new file mode 100644 index 00000000..02512e03 --- /dev/null +++ b/neural_lam/create_graph_with_wmg.py @@ -0,0 +1,209 @@ +"""Create neural-lam graphs by delegating construction to weather-model-graphs. + +Builds the g2m/m2m/m2g graph components with weather-model-graphs (wmg) and +saves them to disk in neural-lam's tensor-on-disk format, replacing the +duplicated logic in ``create_graph.py``. +""" + +# Standard library +import os +from argparse import ArgumentDefaultsHelpFormatter, ArgumentParser + +# Third-party +import numpy as np +import weather_model_graphs as wmg +from loguru import logger + +# Local +from .config import load_config_and_datastore +from .datastore.base import BaseRegularGridDatastore + +ARCHETYPE_FUNCTIONS = { + "keisler": wmg.create.archetype.create_keisler_graph, + "graphcast": wmg.create.archetype.create_graphcast_graph, + "hierarchical": wmg.create.archetype.create_oskarsson_hierarchical_graph, +} + + +def _estimate_grid_node_distance(xy): + """Estimate the average grid node distance from grid coordinates. + + Parameters + ---------- + xy : np.ndarray + Grid coordinates of shape ``(N, 2)``. + + Returns + ------- + float + Estimated average grid node distance in coordinate units. + """ + x_range = np.ptp(xy[:, 0]) + y_range = np.ptp(xy[:, 1]) + n_points = len(xy) + # avg grid node distance ≈ sqrt(area / n_points) + return float(np.sqrt(x_range * y_range / n_points)) + + +def create_graph_from_datastore( + datastore, + output_root_path, + archetype="keisler", + mesh_node_distance=None, + mesh_grid_distance_ratio=3.0, + level_refinement_factor=3, + max_num_levels=None, +): + """Create graph using weather-model-graphs and save in neural-lam format. + + Parameters + ---------- + datastore : BaseRegularGridDatastore + Datastore providing grid coordinates. + output_root_path : str + Directory where the .pt graph files will be saved. + archetype : str + Graph archetype to create: ``"keisler"``, ``"graphcast"``, or + ``"hierarchical"``. + mesh_node_distance : float or None + Distance between created mesh nodes (in coordinate units). If None, + the grid node distance is estimated automatically from the grid + coordinates and multiplied by ``mesh_grid_distance_ratio``. + mesh_grid_distance_ratio : float + Ratio of mesh node distance to grid node distance. Only used when + ``mesh_node_distance`` is None. Default is 3.0. + level_refinement_factor : int + Refinement factor between mesh hierarchy levels. Only used for + ``"graphcast"`` and ``"hierarchical"`` archetypes. + max_num_levels : int or None + Maximum number of mesh hierarchy levels. Only used for ``"graphcast"`` + and ``"hierarchical"`` archetypes. + """ + if not isinstance(datastore, BaseRegularGridDatastore): + raise NotImplementedError( + "Only graph creation for BaseRegularGridDatastore is supported" + ) + + if archetype not in ARCHETYPE_FUNCTIONS: + raise ValueError( + f"Unknown archetype '{archetype}'. " + f"Must be one of: {list(ARCHETYPE_FUNCTIONS.keys())}" + ) + + xy = datastore.get_xy(category="state", stacked=True) + xy = np.array(xy) + + if mesh_node_distance is None: + grid_node_distance = _estimate_grid_node_distance(xy) + mesh_node_distance = grid_node_distance * mesh_grid_distance_ratio + logger.info( + f"mesh_node_distance not given; estimated grid node distance " + f"{grid_node_distance:.2f} x mesh_grid_distance_ratio " + f"{mesh_grid_distance_ratio} -> mesh_node_distance " + f"{mesh_node_distance:.2f}" + ) + + # Build keyword arguments for the archetype function. + # return_components=True is required because + # wmg.save.to_torch_tensors_on_disk() expects the graph as + # separate g2m, m2g and m2m sub-graph components + # rather than a single merged graph. + archetype_kwargs = dict( + coords=xy, + mesh_node_distance=mesh_node_distance, + return_components=True, + ) + + # Only multiscale/hierarchical archetypes accept these parameters + if archetype in ("graphcast", "hierarchical"): + archetype_kwargs["level_refinement_factor"] = level_refinement_factor + archetype_kwargs["max_num_levels"] = max_num_levels + + archetype_fn = ARCHETYPE_FUNCTIONS[archetype] + graph_components = archetype_fn(**archetype_kwargs) + + hierarchical = archetype == "hierarchical" + + wmg.save.to_torch_tensors_on_disk( + graph_components=graph_components, + output_directory=output_root_path, + hierarchical=hierarchical, + ) + + +def cli(input_args=None): + """Command-line interface for graph creation using weather-model-graphs.""" + parser = ArgumentParser( + description="Graph generation for neural-lam using " + "weather-model-graphs (wmg)", + formatter_class=ArgumentDefaultsHelpFormatter, + ) + parser.add_argument( + "--config_path", + type=str, + help="Path to neural-lam configuration file", + ) + parser.add_argument( + "--name", + type=str, + default="multiscale", + help="Name to save graph as (used as subdirectory name)", + ) + parser.add_argument( + "--archetype", + type=str, + default="keisler", + choices=["keisler", "graphcast", "hierarchical"], + help="Graph archetype to create", + ) + parser.add_argument( + "--mesh_node_distance", + type=float, + default=None, + help="Distance between mesh nodes (in coordinate units). " + "If not set, estimated automatically from the grid node distance " + "and --mesh_grid_distance_ratio.", + ) + parser.add_argument( + "--mesh_grid_distance_ratio", + type=float, + default=3.0, + help="Ratio of mesh node distance to grid node distance. " + "Only used when --mesh_node_distance is not set.", + ) + parser.add_argument( + "--level_refinement_factor", + type=int, + default=3, + help="Refinement factor between mesh hierarchy levels " + "(only used for graphcast and hierarchical)", + ) + parser.add_argument( + "--max_num_levels", + type=int, + default=None, + help="Maximum number of mesh levels " + "(only used for graphcast and hierarchical)", + ) + args = parser.parse_args(input_args) + + assert ( + args.config_path is not None + ), "Specify your config with --config_path" + + # Load neural-lam configuration and datastore to use + _, datastore = load_config_and_datastore(config_path=args.config_path) + + create_graph_from_datastore( + datastore=datastore, + output_root_path=os.path.join(datastore.root_path, "graph", args.name), + archetype=args.archetype, + mesh_node_distance=args.mesh_node_distance, + mesh_grid_distance_ratio=args.mesh_grid_distance_ratio, + level_refinement_factor=args.level_refinement_factor, + max_num_levels=args.max_num_levels, + ) + + +if __name__ == "__main__": + cli() diff --git a/pyproject.toml b/pyproject.toml index 113329e1..7485bbeb 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -38,6 +38,7 @@ dependencies = [ "boto3>=1.35.32", "nvidia-ml-py>=13.580.82", "pillow>=9.0.0", + "weather-model-graphs>=0.4.0", ] requires-python = ">=3.10" @@ -50,6 +51,9 @@ cpu = ["torch>=2.12,<2.13"] gpu = ["torch>=2.12,<2.13"] # CUDA 13.0, default GPU build gpu-cu128 = ["torch>=2.11,<2.12"] # CUDA 12.8, last torch series with cu128 wheels +[project.scripts] +create_graph_with_wmg = "neural_lam.create_graph_with_wmg:cli" + [dependency-groups] dev = ["pre-commit>=3.8.0", "pytest>=8.3.2", "pooch>=1.8.2"] diff --git a/tests/test_clamping.py b/tests/test_clamping.py index 8c44b568..bdd148f4 100644 --- a/tests/test_clamping.py +++ b/tests/test_clamping.py @@ -6,7 +6,7 @@ # First-party from neural_lam import config as nlconfig -from neural_lam.create_graph import create_graph_from_datastore +from neural_lam.create_graph_with_wmg import create_graph_from_datastore from neural_lam.datastore.mdp import MDPDatastore from neural_lam.models import GraphLAM from tests.conftest import init_datastore_example @@ -23,7 +23,7 @@ def test_clamping(): create_graph_from_datastore( datastore=datastore, output_root_path=str(graph_dir_path), - n_max_levels=1, + archetype="keisler", ) class ModelArgs: diff --git a/tests/test_datasets.py b/tests/test_datasets.py index 4b35840e..01148ac4 100644 --- a/tests/test_datasets.py +++ b/tests/test_datasets.py @@ -9,7 +9,7 @@ # First-party from neural_lam import config as nlconfig -from neural_lam.create_graph import create_graph_from_datastore +from neural_lam.create_graph_with_wmg import create_graph_from_datastore from neural_lam.datastore import DATASTORES from neural_lam.datastore.base import BaseRegularGridDatastore from neural_lam.models import ForecasterModule @@ -200,7 +200,7 @@ def _create_graph(): create_graph_from_datastore( datastore=datastore, output_root_path=str(graph_dir_path), - n_max_levels=1, + archetype="keisler", ) if not isinstance(datastore, BaseRegularGridDatastore): diff --git a/tests/test_graph_creation.py b/tests/test_graph_creation.py index c97bd191..0b72b66b 100644 --- a/tests/test_graph_creation.py +++ b/tests/test_graph_creation.py @@ -1,6 +1,7 @@ # Standard library import importlib.util import tempfile +import warnings from pathlib import Path # Third-party @@ -13,6 +14,9 @@ METAINFO_FILENAME, create_graph_from_datastore, ) +from neural_lam.create_graph_with_wmg import ( + create_graph_from_datastore as wmg_create_graph_from_datastore, +) from neural_lam.datastore import DATASTORES from neural_lam.datastore.base import BaseRegularGridDatastore from neural_lam.utils import BufferList, load_graph @@ -341,3 +345,70 @@ def test_buffer_list_iter(buffer_list_five): """Iteration yields all buffers in order.""" values = [t.item() for t in buffer_list_five] assert values == [0.0, 1.0, 2.0, 3.0, 4.0] + + +@pytest.mark.parametrize("archetype", ["keisler", "graphcast", "hierarchical"]) +@pytest.mark.parametrize("datastore_name", DATASTORES.keys()) +def test_wmg_graph_creation(datastore_name, archetype): + """Check that graph creation via weather-model-graphs produces a graph + that conforms to the graph storage specification.""" + datastore = init_datastore_example(datastore_name) + + if not isinstance(datastore, BaseRegularGridDatastore): + pytest.skip( + f"Skipping test for {datastore_name} as it is not a regular " + "grid datastore." + ) + + hierarchical = archetype == "hierarchical" + + with tempfile.TemporaryDirectory() as tmpdir: + graph_dir_path = Path(tmpdir) / "graph" / archetype + + wmg_create_graph_from_datastore( + datastore=datastore, + output_root_path=str(graph_dir_path), + archetype=archetype, + ) + + # Validate the wmg-created graph on disk against the graph-storage + # spec and validator introduced in #323. This is the end-to-end + # contract check: a graph built through create_graph_with_wmg (using + # weather-model-graphs' to_torch_tensors_on_disk) must pass the same + # validator neural-lam ships for the on-disk graph format. It covers + # file presence, spec version, container types, edge-index shapes and + # feature dimensions, so those are not re-checked here. + validator = _load_validator_module() + report, _, _ = validator.validate_graph_directory(graph_dir_path) + assert not report.has_fails(), report.summarize() + + # The validator infers whether a graph is hierarchical from its + # contents, so it cannot tell whether the requested archetype was + # honoured. Check that separately. + for file_name in ( + "mesh_up_edge_index.pt", + "mesh_down_edge_index.pt", + "mesh_up_features.pt", + "mesh_down_features.pt", + ): + assert (graph_dir_path / file_name).exists() == hierarchical + + +@pytest.mark.parametrize("datastore_name", DATASTORES.keys()) +def test_old_create_graph_deprecation_warning(datastore_name): + """Check that the old create_graph CLI emits a deprecation warning.""" + # First-party + from neural_lam.create_graph import cli as old_cli + + with warnings.catch_warnings(record=True) as w: + warnings.simplefilter("always") + try: + old_cli(["--config_path", "nonexistent.yaml"]) + except Exception: + pass # We only care about the warning, not the error + + deprecation_warnings = [ + x for x in w if issubclass(x.category, DeprecationWarning) + ] + assert len(deprecation_warnings) >= 1 + assert "create_graph_with_wmg" in str(deprecation_warnings[0].message) diff --git a/tests/test_plot_graph.py b/tests/test_plot_graph.py index 87bd3a3f..7a66f5a5 100644 --- a/tests/test_plot_graph.py +++ b/tests/test_plot_graph.py @@ -7,7 +7,7 @@ # First-party from neural_lam import utils -from neural_lam.create_graph import create_graph_from_datastore +from neural_lam.create_graph_with_wmg import create_graph_from_datastore from neural_lam.plot_graph import ( plot_graph, ) @@ -18,8 +18,9 @@ def graph_fixture(request, tmp_path_factory): """Create a graph from a DummyDatastore and load it back. - Parametrized over graph types: 1level (flat), multiscale (flat multi-level), - and hierarchical. + Parametrized over graph types: 1level (flat, keisler archetype), + multiscale (flat multi-level, graphcast archetype) and hierarchical + (multi-level with up/down edges). Returns ------- @@ -30,14 +31,14 @@ def graph_fixture(request, tmp_path_factory): datastore = DummyDatastore() if graph_name == "hierarchical": - hierarchical = True - n_max_levels = 3 + archetype = "hierarchical" + max_num_levels = 3 elif graph_name == "multiscale": - hierarchical = False - n_max_levels = 3 + archetype = "graphcast" + max_num_levels = 3 elif graph_name == "1level": - hierarchical = False - n_max_levels = 1 + archetype = "keisler" + max_num_levels = None else: raise ValueError(f"Unknown graph_name: {graph_name}") @@ -45,8 +46,8 @@ def graph_fixture(request, tmp_path_factory): create_graph_from_datastore( datastore=datastore, output_root_path=str(graph_dir_path), - hierarchical=hierarchical, - n_max_levels=n_max_levels, + archetype=archetype, + max_num_levels=max_num_levels, ) grid_xy_extent = datastore.get_xy_extent(category="state") diff --git a/tests/test_training.py b/tests/test_training.py index bf1a5884..5ef77522 100644 --- a/tests/test_training.py +++ b/tests/test_training.py @@ -10,7 +10,7 @@ # First-party from neural_lam import config as nlconfig -from neural_lam.create_graph import create_graph_from_datastore +from neural_lam.create_graph_with_wmg import create_graph_from_datastore from neural_lam.datastore import DATASTORES from neural_lam.datastore.base import BaseRegularGridDatastore from neural_lam.models import ForecasterModule @@ -85,7 +85,7 @@ def run_simple_training( create_graph_from_datastore( datastore=datastore, output_root_path=str(graph_dir_path), - n_max_levels=1, + archetype="keisler", ) data_module = WeatherDataModule( diff --git a/uv.lock b/uv.lock index b9d5a7ed..4e18287f 100644 --- a/uv.lock +++ b/uv.lock @@ -2612,6 +2612,7 @@ dependencies = [ { name = "torch-geometric" }, { name = "tueplots" }, { name = "wandb" }, + { name = "weather-model-graphs" }, ] [package.optional-dependencies] @@ -2658,6 +2659,7 @@ requires-dist = [ { name = "torch-geometric", specifier = "==2.3.1" }, { name = "tueplots", specifier = ">=0.0.8" }, { name = "wandb", specifier = ">=0.13.10" }, + { name = "weather-model-graphs", specifier = ">=0.4.0" }, ] provides-extras = ["cpu", "gpu", "gpu-cu128"] @@ -5565,6 +5567,25 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/68/5a/199c59e0a824a3db2b89c5d2dade7ab5f9624dbf6448dc291b46d5ec94d3/wcwidth-0.6.0-py3-none-any.whl", hash = "sha256:1a3a1e510b553315f8e146c54764f4fb6264ffad731b3d78088cdb1478ffbdad", size = 94189, upload-time = "2026-02-06T19:19:39.646Z" }, ] +[[package]] +name = "weather-model-graphs" +version = "0.4.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "loguru" }, + { name = "networkx", version = "3.4.2", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11' or (extra == 'extra-10-neural-lam-cpu' and extra == 'extra-10-neural-lam-gpu') or (extra == 'extra-10-neural-lam-cpu' and extra == 'extra-10-neural-lam-gpu-cu128') or (extra == 'extra-10-neural-lam-gpu' and extra == 'extra-10-neural-lam-gpu-cu128')" }, + { name = "networkx", version = "3.6.1", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.11' or (extra == 'extra-10-neural-lam-cpu' and extra == 'extra-10-neural-lam-gpu') or (extra == 'extra-10-neural-lam-cpu' and extra == 'extra-10-neural-lam-gpu-cu128') or (extra == 'extra-10-neural-lam-gpu' and extra == 'extra-10-neural-lam-gpu-cu128')" }, + { name = "numpy", version = "2.2.6", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11' or (extra == 'extra-10-neural-lam-cpu' and extra == 'extra-10-neural-lam-gpu') or (extra == 'extra-10-neural-lam-cpu' and extra == 'extra-10-neural-lam-gpu-cu128') or (extra == 'extra-10-neural-lam-gpu' and extra == 'extra-10-neural-lam-gpu-cu128')" }, + { name = "numpy", version = "2.4.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.11' or (extra == 'extra-10-neural-lam-cpu' and extra == 'extra-10-neural-lam-gpu') or (extra == 'extra-10-neural-lam-cpu' and extra == 'extra-10-neural-lam-gpu-cu128') or (extra == 'extra-10-neural-lam-gpu' and extra == 'extra-10-neural-lam-gpu-cu128')" }, + { name = "pyproj", version = "3.7.1", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11' or (extra == 'extra-10-neural-lam-cpu' and extra == 'extra-10-neural-lam-gpu') or (extra == 'extra-10-neural-lam-cpu' and extra == 'extra-10-neural-lam-gpu-cu128') or (extra == 'extra-10-neural-lam-gpu' and extra == 'extra-10-neural-lam-gpu-cu128')" }, + { name = "pyproj", version = "3.7.2", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.11' or (extra == 'extra-10-neural-lam-cpu' and extra == 'extra-10-neural-lam-gpu') or (extra == 'extra-10-neural-lam-cpu' and extra == 'extra-10-neural-lam-gpu-cu128') or (extra == 'extra-10-neural-lam-gpu' and extra == 'extra-10-neural-lam-gpu-cu128')" }, + { name = "scipy", version = "1.15.3", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11' or (extra == 'extra-10-neural-lam-cpu' and extra == 'extra-10-neural-lam-gpu') or (extra == 'extra-10-neural-lam-cpu' and extra == 'extra-10-neural-lam-gpu-cu128') or (extra == 'extra-10-neural-lam-gpu' and extra == 'extra-10-neural-lam-gpu-cu128')" }, + { name = "scipy", version = "1.17.1", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.11' or (extra == 'extra-10-neural-lam-cpu' and extra == 'extra-10-neural-lam-gpu') or (extra == 'extra-10-neural-lam-cpu' and extra == 'extra-10-neural-lam-gpu-cu128') or (extra == 'extra-10-neural-lam-gpu' and extra == 'extra-10-neural-lam-gpu-cu128')" }, +] +wheels = [ + { url = "https://files.pythonhosted.org/packages/10/fc/de383b841fb9fe75960d55102fd16108ea5d3b245471956eac2ec1899eee/weather_model_graphs-0.4.0-py3-none-any.whl", hash = "sha256:261cee5d72bac06b6f36f8e3d2f5c7850ba56976d5759923710353d9459639dc", size = 44677, upload-time = "2026-07-28T09:23:59.658Z" }, +] + [[package]] name = "werkzeug" version = "3.1.7"