-
Notifications
You must be signed in to change notification settings - Fork 275
Add create_graph_with_wmg.py CLI using weather-model-graphs #596
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from all commits
616583a
fbe388e
e514831
35fee71
4f27342
deff22b
3ee144e
1d90155
7cf64bc
bba00b3
2e863da
eaa9865
8665d28
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -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, | ||
|
leifdenby marked this conversation as resolved.
|
||
| ) | ||
|
|
||
| # 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", | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. I think it is a bit confusing that we use both term "spacing" and "distance" for the mesh/grid resolution. I appreciate that wmg uses "distance", so maybe should use the same here? "resolution" would probably be an even better term. What do you think?
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Yeah, Leif - that's me mixing the two, not wmg. I had a look and there's no convention to protect either way: "spacing" doesn't appear anywhere in the codebase outside this file, and "resolution" only turns up in a comment in I'd go with "distance". wmg's parameter is That would give One thing I'd like your view on:
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
I agree - good point!
Very well caught! Nice, yes rename the arg
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Renamed in 8665d28, both the parameter and the CLI flag: |
||
| 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() | ||
Uh oh!
There was an error while loading. Please reload this page.