diff --git a/examples/weathernext.py b/examples/weathernext.py new file mode 100644 index 000000000..7948dae56 --- /dev/null +++ b/examples/weathernext.py @@ -0,0 +1,241 @@ +#!/usr/bin/env python +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +"""Build and run a WeatherNext-style one-step forecast ONNX model. + +The formal Mobius support lives in ``mobius.models.WeatherNextModel``, +``mobius.tasks.WeatherNextForecastTask``, and ``mobius.integrations.weathernext``. +This script demonstrates that path and can run the exported model on either: + +* an ``.npz`` file with ``input_state``, ``forcings``, and ``sample_noise`` arrays, or +* a local xarray NetCDF/Zarr weather dataset plus selected variable names. + +If no data or checkpoint is provided, the script uses deterministic demo weights +and synthetic inputs so the ONNX workflow remains runnable in a fresh checkout. + +Usage:: + + PYTHONPATH=src python examples/weathernext.py output/weathernext-mini --validate + + PYTHONPATH=src python examples/weathernext.py output/weathernext-era5 \ + --input-data era5_sample.npz --weights converted_weathernext_weights.npz --run + + PYTHONPATH=src python examples/weathernext.py output/weathernext-xarray \ + --input-data weatherbench_sample.zarr \ + --input-variable-names 2m_temperature mean_sea_level_pressure \ + --forcing-variable-names toa_incident_solar_radiation --run +""" + +from __future__ import annotations + +import argparse +import os +import sys + +import numpy as np +import onnx_ir as ir + +from mobius import WeatherNextConfig +from mobius.integrations.weathernext import ( + build_weathernext_package, + infer_config_from_feeds, + load_npz_forecast_inputs, + load_npz_weights, + load_xarray_forecast_inputs, +) + + +def _resolve_dtype(name: str) -> ir.DataType: + if name == "f32": + return ir.DataType.FLOAT + if name == "f16": + return ir.DataType.FLOAT16 + raise ValueError(f"Unsupported dtype: {name}") + + +def _load_real_data(args: argparse.Namespace) -> dict[str, np.ndarray] | None: + if args.input_data is None: + return None + if args.input_data.endswith(".npz"): + return load_npz_forecast_inputs(args.input_data) + if not args.input_variable_names or not args.forcing_variable_names: + raise ValueError( + "--input-variable-names and --forcing-variable-names are required for xarray data" + ) + return load_xarray_forecast_inputs( + args.input_data, + input_variables=args.input_variable_names, + forcing_variables=args.forcing_variable_names, + noise_channels=args.noise_channels, + batch_index=args.batch_index, + sample_noise_seed=args.sample_noise_seed, + ) + + +def _synthetic_feeds(config: WeatherNextConfig) -> dict[str, np.ndarray]: + rng = np.random.default_rng(42) + return { + "input_state": rng.standard_normal( + (1, config.lat, config.lon, config.input_variables) + ).astype(np.float32), + "forcings": rng.standard_normal( + (1, config.lat, config.lon, config.forcing_variables) + ).astype(np.float32), + "sample_noise": rng.standard_normal( + (1, config.lat, config.lon, config.noise_channels) + ).astype(np.float32), + } + + +def _numpy_dtype(dtype: ir.DataType) -> type[np.float32 | np.float16]: + if dtype == ir.DataType.FLOAT: + return np.float32 + if dtype == ir.DataType.FLOAT16: + return np.float16 + raise ValueError(f"Unsupported WeatherNext feed dtype: {dtype}") + + +def _cast_feeds_to_dtype( + feeds: dict[str, np.ndarray], dtype: ir.DataType +) -> dict[str, np.ndarray]: + feed_dtype = _numpy_dtype(dtype) + return {name: np.asarray(value, dtype=feed_dtype) for name, value in feeds.items()} + + +def _config_from_args( + args: argparse.Namespace, feeds: dict[str, np.ndarray] | None +) -> WeatherNextConfig: + dtype = _resolve_dtype(args.dtype) + if feeds is not None: + return infer_config_from_feeds( + feeds, + mesh_nodes=args.mesh_nodes, + hidden_size=args.hidden_size, + output_variables=args.output_variables, + intermediate_size=args.intermediate_size, + num_hidden_layers=args.num_hidden_layers, + dtype=dtype, + ) + return WeatherNextConfig( + lat=args.lat, + lon=args.lon, + mesh_nodes=args.mesh_nodes, + input_variables=args.input_variables, + forcing_variables=args.forcing_variables, + noise_channels=args.noise_channels, + output_variables=args.input_variables + if args.output_variables is None + else args.output_variables, + hidden_size=args.hidden_size, + intermediate_size=args.intermediate_size or 4 * args.hidden_size, + num_hidden_layers=args.num_hidden_layers, + dtype=dtype, + ) + + +def _run_with_ort(output_dir: str, feeds: dict[str, np.ndarray], dtype: ir.DataType) -> None: + try: + import onnxruntime as ort + except ImportError: + print("onnxruntime is not installed; skipping inference.", file=sys.stderr) + return + + model_path = os.path.join(output_dir, "model.onnx") + sess = ort.InferenceSession(model_path, providers=["CPUExecutionProvider"]) + (next_state,) = sess.run(["next_state"], _cast_feeds_to_dtype(feeds, dtype)) + print(f"Inference output next_state shape: {next_state.shape}") + print(f"Inference output range: [{next_state.min():.6f}, {next_state.max():.6f}]") + + +def _parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser( + description="Build a Mobius WeatherNext one-step forecast ONNX graph.", + ) + parser.add_argument("output_dir", help="Directory to save model.onnx and model.onnx.data.") + parser.add_argument("--lat", type=int, default=4, help="Synthetic-data latitude points.") + parser.add_argument("--lon", type=int, default=8, help="Synthetic-data longitude points.") + parser.add_argument("--mesh-nodes", type=int, default=6, help="Forecast mesh nodes.") + parser.add_argument( + "--input-variables", type=int, default=5, help="Synthetic input channels." + ) + parser.add_argument( + "--forcing-variables", type=int, default=2, help="Synthetic forcing channels." + ) + parser.add_argument( + "--noise-channels", type=int, default=2, help="Stochastic noise channels." + ) + parser.add_argument( + "--output-variables", + type=int, + help="Output channels. Defaults to input-variable count for synthetic and real data.", + ) + parser.add_argument("--hidden-size", type=int, default=16, help="Latent feature size.") + parser.add_argument("--intermediate-size", type=int, help="Mesh MLP intermediate size.") + parser.add_argument("--num-hidden-layers", type=int, default=1, help="Mesh update blocks.") + parser.add_argument( + "--dtype", choices=["f32", "f16"], default="f32", help="ONNX weight dtype." + ) + parser.add_argument("--weights", help="Optional Mobius-aligned WeatherNext weights .npz.") + parser.add_argument( + "--input-data", + help="Optional .npz, NetCDF, or Zarr weather sample used for real-data inference.", + ) + parser.add_argument( + "--input-variable-names", + nargs="+", + help="xarray variables stacked into input_state channels.", + ) + parser.add_argument( + "--forcing-variable-names", + nargs="+", + help="xarray variables stacked into forcings channels.", + ) + parser.add_argument( + "--batch-index", type=int, default=0, help="Time/batch index for xarray." + ) + parser.add_argument( + "--sample-noise-seed", type=int, default=0, help="Generated noise seed." + ) + parser.add_argument("--run", action="store_true", help="Run one ONNX Runtime inference.") + parser.add_argument( + "--validate", + action="store_true", + help="Alias for --run, kept for the original self-contained demo workflow.", + ) + return parser.parse_args() + + +def main() -> None: + args = _parse_args() + feeds = _load_real_data(args) + config = _config_from_args(args, feeds) + config.validate() + + print("Building WeatherNext one-step forecast ONNX graph...") + print(f" grid: {config.lat} x {config.lon} ({config.grid_points} cells)") + print(f" mesh nodes: {config.mesh_nodes}") + print( + " channels: " + f"input={config.input_variables}, forcing={config.forcing_variables}, " + f"noise={config.noise_channels}, output={config.output_variables}" + ) + + weights = load_npz_weights(args.weights) if args.weights else None + package = build_weathernext_package(config, weights=weights) + model = package["model"] + print(f"Built model with {model.graph.num_nodes()} ONNX nodes.") + + package.save(args.output_dir, check_weights=True, progress_bar=False) + print(f"Saved WeatherNext package to {args.output_dir!r}.") + + if args.run or args.validate: + _run_with_ort( + args.output_dir, + feeds if feeds is not None else _synthetic_feeds(config), + config.dtype, + ) + + +if __name__ == "__main__": + main() diff --git a/src/mobius/__init__.py b/src/mobius/__init__.py index 0962f505a..373e24020 100644 --- a/src/mobius/__init__.py +++ b/src/mobius/__init__.py @@ -30,6 +30,9 @@ "VisionConfig", "VisionLanguageConfig", "WhisperConfig", + "WeatherNextConfig", + "WeatherNextForecastTask", + "WeatherNextModel", "WorldModelConfig", "WorldModelTask", "YolosConfig", @@ -78,6 +81,7 @@ SegformerConfig, VisionConfig, VisionLanguageConfig, + WeatherNextConfig, WhisperConfig, WorldModelConfig, YolosConfig, @@ -95,5 +99,5 @@ from mobius._weight_loading import apply_weights from mobius.integrations.gguf import build_from_gguf from mobius.integrations.nemo import build_from_nemo -from mobius.models import MLPWorldModel -from mobius.tasks import CausalLMTask, ModelTask, WorldModelTask +from mobius.models import MLPWorldModel, WeatherNextModel +from mobius.tasks import CausalLMTask, ModelTask, WeatherNextForecastTask, WorldModelTask diff --git a/src/mobius/_configs/__init__.py b/src/mobius/_configs/__init__.py index afb241a16..d5e30e945 100644 --- a/src/mobius/_configs/__init__.py +++ b/src/mobius/_configs/__init__.py @@ -79,6 +79,7 @@ TTSConfig, VisionConfig, ) +from mobius._configs._weathernext import WeatherNextConfig from mobius._configs._world_model import WorldModelConfig __all__ = [ @@ -118,6 +119,7 @@ "VisionConfig", "VisionLanguageConfig", "WhisperConfig", + "WeatherNextConfig", "WorldModelConfig", "YolosConfig", "Zamba2Config", diff --git a/src/mobius/_configs/_weathernext.py b/src/mobius/_configs/_weathernext.py new file mode 100644 index 000000000..f38db7149 --- /dev/null +++ b/src/mobius/_configs/_weathernext.py @@ -0,0 +1,62 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +"""Configuration for WeatherNext-style one-step forecast graphs.""" + +from __future__ import annotations + +import dataclasses + +from mobius._configs._base import BaseModelConfig + + +@dataclasses.dataclass +class WeatherNextConfig(BaseModelConfig): + """Configuration for grid→mesh→grid WeatherNext forecast modules. + + Shapes exclude the leading batch dimension. The task exposes a one-step + forecast contract: + ``input_state + forcings + sample_noise -> next_state``. + """ + + lat: int = 4 + lon: int = 8 + mesh_nodes: int = 6 + input_variables: int = 5 + forcing_variables: int = 2 + noise_channels: int = 2 + output_variables: int = 5 + hidden_size: int = 16 + intermediate_size: int = 64 + num_hidden_layers: int = 1 + hidden_act: str | None = "silu" + + @property + def grid_points(self) -> int: + """Number of lat/lon grid cells.""" + return self.lat * self.lon + + @property + def encoder_channels(self) -> int: + """Per-grid-cell input channels after concatenating all inputs.""" + return self.input_variables + self.forcing_variables + self.noise_channels + + def validate(self) -> None: + """Validate dimensions required by the WeatherNext forecast task.""" + for name in ( + "lat", + "lon", + "mesh_nodes", + "input_variables", + "forcing_variables", + "noise_channels", + "output_variables", + "hidden_size", + "intermediate_size", + "num_hidden_layers", + ): + value = getattr(self, name) + if not isinstance(value, int) or isinstance(value, bool) or value <= 0: + raise ValueError(f"{name} must be a positive integer") + if self.hidden_act is None: + raise ValueError("hidden_act must be set") diff --git a/src/mobius/integrations/weathernext.py b/src/mobius/integrations/weathernext.py new file mode 100644 index 000000000..52a440a41 --- /dev/null +++ b/src/mobius/integrations/weathernext.py @@ -0,0 +1,211 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +"""Helpers for building and running WeatherNext-style Mobius packages.""" + +from __future__ import annotations + +from pathlib import Path +from typing import Any + +import numpy as np +import onnx_ir as ir +import torch + +from mobius._builder import build_from_module +from mobius._configs import WeatherNextConfig +from mobius._model_package import ModelPackage +from mobius.models import WeatherNextModel + +_INPUT_NAMES = ("input_state", "forcings", "sample_noise") + + +def load_npz_weights(path: str | Path) -> dict[str, torch.Tensor]: + """Load a Mobius-aligned WeatherNext state dict from an ``.npz`` file.""" + with np.load(path) as data: + return {name: torch.from_numpy(np.asarray(data[name])) for name in data.files} + + +def create_demo_state_dict( + config: WeatherNextConfig, + *, + seed: int = 20260810, +) -> dict[str, torch.Tensor]: + """Create deterministic weights for examples and tests. + + These weights are intentionally small and are not a trained WeatherNext + checkpoint. Pass ``weights=load_npz_weights(...)`` to + :func:`build_weathernext_package` for converted checkpoint weights. + """ + rng = np.random.default_rng(seed) + + def parameter(shape: tuple[int, ...]) -> torch.Tensor: + return torch.from_numpy(rng.standard_normal(shape).astype(np.float32) * 0.05) + + state: dict[str, torch.Tensor] = { + "grid_encoder.weight": parameter((config.hidden_size, config.encoder_channels)), + "grid_encoder.bias": parameter((config.hidden_size,)), + "grid_decoder.weight": parameter((config.output_variables, config.hidden_size)), + "grid_decoder.bias": parameter((config.output_variables,)), + } + for layer_idx in range(config.num_hidden_layers): + state[f"mesh_update_in.{layer_idx}.weight"] = parameter( + (config.intermediate_size, config.hidden_size) + ) + state[f"mesh_update_in.{layer_idx}.bias"] = parameter((config.intermediate_size,)) + state[f"mesh_update_out.{layer_idx}.weight"] = parameter( + (config.hidden_size, config.intermediate_size) + ) + state[f"mesh_update_out.{layer_idx}.bias"] = parameter((config.hidden_size,)) + return state + + +def build_weathernext_package( + config: WeatherNextConfig, + *, + weights: dict[str, torch.Tensor] | None = None, + execution_provider: str = "default", +) -> ModelPackage: + """Build a WeatherNext one-step forecast package and apply weights.""" + package = build_from_module( + WeatherNextModel(config), + config, + task="weathernext-forecast", + execution_provider=execution_provider, + ) + package.apply_weights(weights if weights is not None else create_demo_state_dict(config)) + return package + + +def load_npz_forecast_inputs(path: str | Path) -> dict[str, np.ndarray]: + """Load ``input_state``, ``forcings``, and ``sample_noise`` arrays from ``.npz``.""" + with np.load(path) as data: + missing = [name for name in _INPUT_NAMES if name not in data] + if missing: + expected = ", ".join(_INPUT_NAMES) + missing_names = ", ".join(missing) + raise ValueError( + f"WeatherNext input file is missing {missing_names}; expected keys: {expected}" + ) + feeds = {name: np.asarray(data[name], dtype=np.float32) for name in _INPUT_NAMES} + return feeds + + +def load_xarray_forecast_inputs( + path: str | Path, + *, + input_variables: list[str], + forcing_variables: list[str], + noise_channels: int, + batch_index: int = 0, + sample_noise_seed: int = 0, +) -> dict[str, np.ndarray]: + """Load WeatherNext inputs from a local xarray NetCDF or Zarr dataset. + + The selected variables must share latitude/longitude dimensions. Optional + time dimensions are indexed by ``batch_index`` and each variable is stacked + into the trailing channel dimension expected by the ONNX graph. The default + noise seed is deterministic so examples are reproducible; pass a different + seed when sampling stochastic forecast noise for real workflows. + """ + try: + import xarray as xr + except ImportError as e: # pragma: no cover - exercised only without optional xarray + raise ImportError("xarray is required to read NetCDF/Zarr WeatherNext inputs") from e + + path = Path(path) + dataset = ( + xr.open_zarr(path) + if path.is_dir() or path.suffix == ".zarr" + else xr.open_dataset(path) + ) + try: + input_state = _stack_xarray_variables( + dataset, + input_variables, + batch_index=batch_index, + ) + forcings = _stack_xarray_variables( + dataset, + forcing_variables, + batch_index=batch_index, + ) + # The seed controls only the generated stochastic noise input. + rng = np.random.default_rng(sample_noise_seed) + sample_noise = rng.standard_normal( + (input_state.shape[0], input_state.shape[1], input_state.shape[2], noise_channels) + ).astype(np.float32) + return { + "input_state": input_state, + "forcings": forcings, + "sample_noise": sample_noise, + } + finally: + dataset.close() + + +def infer_config_from_feeds( + feeds: dict[str, np.ndarray], + *, + mesh_nodes: int, + hidden_size: int, + output_variables: int | None = None, + intermediate_size: int | None = None, + num_hidden_layers: int = 1, + dtype: ir.DataType = ir.DataType.FLOAT, +) -> WeatherNextConfig: + """Infer a WeatherNext config from loaded forecast input arrays.""" + input_state = feeds["input_state"] + forcings = feeds["forcings"] + sample_noise = feeds["sample_noise"] + if input_state.ndim != 4 or forcings.ndim != 4 or sample_noise.ndim != 4: + raise ValueError( + "WeatherNext inputs must be rank-4 [batch, lat, lon, channels] arrays" + ) + if ( + input_state.shape[:3] != forcings.shape[:3] + or input_state.shape[:3] != sample_noise.shape[:3] + ): + raise ValueError("WeatherNext inputs must have matching batch/lat/lon dimensions") + return WeatherNextConfig( + lat=int(input_state.shape[1]), + lon=int(input_state.shape[2]), + mesh_nodes=mesh_nodes, + input_variables=int(input_state.shape[3]), + forcing_variables=int(forcings.shape[3]), + noise_channels=int(sample_noise.shape[3]), + output_variables=int( + input_state.shape[3] if output_variables is None else output_variables + ), + hidden_size=hidden_size, + intermediate_size=intermediate_size or 4 * hidden_size, + num_hidden_layers=num_hidden_layers, + dtype=dtype, + ) + + +def _stack_xarray_variables( + dataset: Any, + names: list[str], + *, + batch_index: int, +) -> np.ndarray: + if not names: + raise ValueError("At least one xarray variable name is required") + arrays = [] + for name in names: + if name not in dataset: + raise KeyError(f"Variable {name!r} not found in dataset") + value = dataset[name] + for dim in value.dims: + if dim.lower() in {"time", "batch"}: + value = value.isel({dim: batch_index}) + array = np.asarray(value, dtype=np.float32) + if array.ndim != 2: + raise ValueError( + f"Variable {name!r} must resolve to a 2-D lat/lon array after " + f"time/batch selection, got shape {array.shape}" + ) + arrays.append(array) + stacked = np.stack(arrays, axis=-1) + return stacked[None, ...] diff --git a/src/mobius/models/__init__.py b/src/mobius/models/__init__.py index 9ca9e2ce3..dad4c3f28 100644 --- a/src/mobius/models/__init__.py +++ b/src/mobius/models/__init__.py @@ -148,6 +148,7 @@ "VideoAutoencoderModel", "Wav2Vec2ForCTCModel", "Wav2Vec2Model", + "WeatherNextModel", "WhisperForConditionalGeneration", "MLPWorldModel", "XLMCausalLMModel", @@ -309,6 +310,7 @@ from mobius.models.vit import ViTModel from mobius.models.wav2vec2 import Wav2Vec2Model from mobius.models.wav2vec2_ctc import Wav2Vec2ForCTCModel +from mobius.models.weathernext import WeatherNextModel from mobius.models.whisper import WhisperForConditionalGeneration from mobius.models.world_model import MLPWorldModel from mobius.models.xlm import XLMCausalLMModel diff --git a/src/mobius/models/weathernext.py b/src/mobius/models/weathernext.py new file mode 100644 index 000000000..78a6e971e --- /dev/null +++ b/src/mobius/models/weathernext.py @@ -0,0 +1,141 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +"""WeatherNext-style grid→mesh→grid forecast model. + +This module provides Mobius-native ONNX graph construction for the one-step +forecast contract used by WeatherNext-family models. It does not trace JAX; +instead, each graph stage is declared directly with ONNX ops. +""" + +from __future__ import annotations + +from typing import TYPE_CHECKING + +import numpy as np +import onnx_ir as ir +from onnxscript import OpBuilder, nn + +from mobius._configs import WeatherNextConfig +from mobius.components import Linear, get_activation + +if TYPE_CHECKING: + import torch + + +def projection_matrix(rows: int, cols: int) -> np.ndarray: + """Create a normalized deterministic projection between grid and mesh points.""" + row_positions = np.linspace(0.0, 1.0, rows, dtype=np.float32)[:, None] + col_positions = np.linspace(0.0, 1.0, cols, dtype=np.float32)[None, :] + distance = np.abs(row_positions - col_positions) + weights = np.maximum(1.0 - 2.0 * distance, 0.0) + weights += 1e-3 + weights /= weights.sum(axis=1, keepdims=True) + return weights.astype(np.float32) + + +class WeatherNextModel(nn.Module): + """One-step WeatherNext-style grid→mesh→grid forecast module. + + Inputs: + - input_state: ``[batch, lat, lon, input_variables]`` + - forcings: ``[batch, lat, lon, forcing_variables]`` + - sample_noise: ``[batch, lat, lon, noise_channels]`` + + Output: + - next_state: ``[batch, lat, lon, output_variables]`` + """ + + default_task = "weathernext-forecast" + config_class = WeatherNextConfig + category = "Weather" + + def __init__(self, config: WeatherNextConfig): + super().__init__() + config.validate() + self.config = config + self.grid_encoder = Linear(config.encoder_channels, config.hidden_size) + self.mesh_update_in = nn.ModuleList( + [ + Linear(config.hidden_size, config.intermediate_size) + for _ in range(config.num_hidden_layers) + ] + ) + self.mesh_update_out = nn.ModuleList( + [ + Linear(config.intermediate_size, config.hidden_size) + for _ in range(config.num_hidden_layers) + ] + ) + self.grid_decoder = Linear(config.hidden_size, config.output_variables) + self._activation = get_activation(config.hidden_act) + + # Fixed topology projections model the grid↔mesh connectivity. A real + # checkpoint may override these names with learned/sparse projection data. + self.grid_to_mesh = nn.Parameter( + [config.mesh_nodes, config.grid_points], + data=ir.tensor(projection_matrix(config.mesh_nodes, config.grid_points)), + ) + self.mesh_to_grid = nn.Parameter( + [config.grid_points, config.mesh_nodes], + data=ir.tensor(projection_matrix(config.grid_points, config.mesh_nodes)), + ) + + def forward( + self, + op: OpBuilder, + input_state: ir.Value, + forcings: ir.Value, + sample_noise: ir.Value, + ) -> ir.Value: + config = self.config + + # Concatenate per-cell weather variables, future forcings, and stochastic + # noise: [B, lat, lon, input+forcing+noise]. + grid_features = op.Concat(input_state, forcings, sample_noise, axis=-1) + + # Encode each lat/lon cell independently: [B, lat, lon, hidden]. + grid_latent = self._activation(op, self.grid_encoder(op, grid_features)) + batch_dim = op.Shape(grid_latent, start=0, end=1) + flat_grid_shape = op.Concat( + batch_dim, + op.Constant(value_ints=[config.grid_points, config.hidden_size]), + axis=0, + ) + # Flatten the spatial grid: [B, lat, lon, hidden] -> [B, grid_points, hidden]. + grid_points = op.Reshape(grid_latent, flat_grid_shape) + + # Aggregate encoded grid cells onto mesh nodes: + # [1, mesh_nodes, grid_points] @ [B, grid_points, hidden] + # -> [B, mesh_nodes, hidden]. + mesh_latent = op.MatMul(op.Unsqueeze(self.grid_to_mesh, [0]), grid_points) + + # Apply one or more residual mesh-update MLP blocks. + for update_in, update_out in zip( + self.mesh_update_in, self.mesh_update_out, strict=True + ): + mesh_delta = self._activation(op, update_in(op, mesh_latent)) + mesh_latent = op.Add(mesh_latent, update_out(op, mesh_delta)) + + # Decode mesh latents back to grid cells and retain the encoded-grid + # residual: [1, grid_points, mesh_nodes] @ [B, mesh_nodes, hidden] + # -> [B, grid_points, hidden]. + grid_delta = op.MatMul(op.Unsqueeze(self.mesh_to_grid, [0]), mesh_latent) + grid_points = op.Add(grid_points, grid_delta) + + # Return one forecast step on the original grid: + # [B, grid_points, output_variables] -> [B, lat, lon, output_variables]. + forecast_points = self.grid_decoder(op, grid_points) + forecast_shape = op.Concat( + batch_dim, + op.Constant(value_ints=[config.lat, config.lon, config.output_variables]), + axis=0, + ) + return op.Reshape(forecast_points, forecast_shape) + + def preprocess_weights( + self, + state_dict: dict[str, torch.Tensor], + ) -> dict[str, torch.Tensor]: + """Return WeatherNext weights unchanged after external integration mapping.""" + return state_dict diff --git a/src/mobius/models/weathernext_test.py b/src/mobius/models/weathernext_test.py new file mode 100644 index 000000000..982022eec --- /dev/null +++ b/src/mobius/models/weathernext_test.py @@ -0,0 +1,132 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +from __future__ import annotations + +import numpy as np +import onnx_ir as ir +import pytest + +from mobius import ( + WeatherNextConfig, + WeatherNextForecastTask, + WeatherNextModel, + build_from_module, +) +from mobius.integrations.weathernext import ( + build_weathernext_package, + create_demo_state_dict, + infer_config_from_feeds, + load_npz_forecast_inputs, +) +from mobius.tasks import TASK_REGISTRY, get_task + + +def _config() -> WeatherNextConfig: + return WeatherNextConfig( + lat=3, + lon=4, + mesh_nodes=5, + input_variables=2, + forcing_variables=1, + noise_channels=1, + output_variables=2, + hidden_size=8, + intermediate_size=16, + num_hidden_layers=2, + ) + + +class TestWeatherNextConfig: + def test_derived_sizes(self): + config = _config() + assert config.grid_points == 12 + assert config.encoder_channels == 4 + + @pytest.mark.parametrize( + ("field", "value"), + [ + ("lat", 0), + ("lon", True), + ("mesh_nodes", -1), + ("input_variables", 0), + ("forcing_variables", 0), + ("noise_channels", 0), + ("output_variables", 0), + ("hidden_size", 0), + ("intermediate_size", 0), + ("num_hidden_layers", 0), + ("hidden_act", None), + ], + ) + def test_invalid_config_raises(self, field, value): + config = _config() + setattr(config, field, value) + with pytest.raises(ValueError): + config.validate() + + +class TestWeatherNextForecastTask: + def test_registered(self): + assert TASK_REGISTRY["weathernext-forecast"] is WeatherNextForecastTask + assert isinstance(get_task("weathernext-forecast"), WeatherNextForecastTask) + + def test_graph_contract(self, tmp_path): + config = _config() + package = build_weathernext_package(config) + package.save(tmp_path, check_weights=True, progress_bar=False) + + model = ir.load(tmp_path / "model.onnx") + assert [value.name for value in model.graph.inputs] == list( + WeatherNextForecastTask.input_names + ) + assert [value.name for value in model.graph.outputs] == list( + WeatherNextForecastTask.output_names + ) + assert model.graph.outputs[0].shape == ir.Shape(["batch", 3, 4, 2]) + assert model.graph.name == "weathernext_one_step_forecast" + + +def test_weathernext_model_requires_weights_for_standard_build(tmp_path): + config = _config() + package = build_from_module(WeatherNextModel(config), config, task="weathernext-forecast") + package.apply_weights(create_demo_state_dict(config)) + package.save(tmp_path, check_weights=True, progress_bar=False) + + +def test_npz_inputs_infer_config(tmp_path): + input_path = tmp_path / "sample.npz" + np.savez( + input_path, + input_state=np.zeros((2, 5, 6, 3), dtype=np.float32), + forcings=np.zeros((2, 5, 6, 2), dtype=np.float32), + sample_noise=np.zeros((2, 5, 6, 1), dtype=np.float32), + ) + + feeds = load_npz_forecast_inputs(input_path) + config = infer_config_from_feeds( + feeds, + mesh_nodes=7, + hidden_size=8, + output_variables=4, + ) + + assert config.lat == 5 + assert config.lon == 6 + assert config.mesh_nodes == 7 + assert config.input_variables == 3 + assert config.forcing_variables == 2 + assert config.noise_channels == 1 + assert config.output_variables == 4 + + +def test_npz_inputs_report_missing_keys(tmp_path): + input_path = tmp_path / "missing.npz" + np.savez( + input_path, + input_state=np.zeros((1, 2, 3, 1), dtype=np.float32), + forcings=np.zeros((1, 2, 3, 1), dtype=np.float32), + ) + + with pytest.raises(ValueError, match="sample_noise"): + load_npz_forecast_inputs(input_path) diff --git a/src/mobius/tasks/__init__.py b/src/mobius/tasks/__init__.py index d7f947ffc..2a780eb89 100644 --- a/src/mobius/tasks/__init__.py +++ b/src/mobius/tasks/__init__.py @@ -67,6 +67,7 @@ "VAETask", "VideoDenoisingTask", "VisionLanguageTask", + "WeatherNextForecastTask", "WorldModelTask", "build_decoder_from_embeds", "build_embedding_from_features", @@ -129,6 +130,7 @@ QwenVLTask, VisionLanguageTask, ) +from mobius.tasks._weathernext import WeatherNextForecastTask from mobius.tasks._world_model import WorldModelTask # --------------------------------------------------------------------------- @@ -181,6 +183,7 @@ "ssm2-text-generation": SSM2CausalLMTask, "tts": TTSTask, "video-denoising": VideoDenoisingTask, + "weathernext-forecast": WeatherNextForecastTask, "world-model": WorldModelTask, } diff --git a/src/mobius/tasks/_weathernext.py b/src/mobius/tasks/_weathernext.py new file mode 100644 index 000000000..719bfe550 --- /dev/null +++ b/src/mobius/tasks/_weathernext.py @@ -0,0 +1,61 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +"""WeatherNext forecast task wiring.""" + +from __future__ import annotations + +from typing import ClassVar + +import onnx_ir as ir +from onnxscript import nn + +from mobius._configs import WeatherNextConfig +from mobius._model_package import ModelPackage +from mobius.tasks._base import ModelTask, _make_graph, _make_model + + +class WeatherNextForecastTask(ModelTask): + """Build a one-step WeatherNext forecast graph. + + Inputs: + - input_state: ``[batch, lat, lon, input_variables]`` + - forcings: ``[batch, lat, lon, forcing_variables]`` + - sample_noise: ``[batch, lat, lon, noise_channels]`` + + Outputs: + - next_state: ``[batch, lat, lon, output_variables]`` + """ + + input_names: ClassVar[tuple[str, ...]] = ("input_state", "forcings", "sample_noise") + output_names: ClassVar[tuple[str, ...]] = ("next_state",) + model_roles: ClassVar[dict[str, str]] = {"model": "forecast"} + + def build( + self, + module: nn.Module, + config: WeatherNextConfig, + ) -> ModelPackage: + config.validate() + batch = ir.SymbolicDim("batch") + graph, builder = _make_graph(name="weathernext_one_step_forecast") + + input_state = builder.input( + self.input_names[0], + dtype=config.dtype, + shape=[batch, config.lat, config.lon, config.input_variables], + ) + forcings = builder.input( + self.input_names[1], + dtype=config.dtype, + shape=[batch, config.lat, config.lon, config.forcing_variables], + ) + sample_noise = builder.input( + self.input_names[2], + dtype=config.dtype, + shape=[batch, config.lat, config.lon, config.noise_channels], + ) + + next_state = module(builder.op, input_state, forcings, sample_noise) + builder.add_output(next_state, self.output_names[0]) + return ModelPackage({"model": _make_model(graph)}, config=config) diff --git a/tests/weathernext_example_test.py b/tests/weathernext_example_test.py new file mode 100644 index 000000000..c141327fa --- /dev/null +++ b/tests/weathernext_example_test.py @@ -0,0 +1,94 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +from __future__ import annotations + +import os +import subprocess +import sys +from pathlib import Path + +import numpy as np +import onnx_ir as ir +import pytest + + +def test_weathernext_example_builds_from_real_npz_inputs(tmp_path): + repo_root = Path(__file__).parents[1] + input_path = tmp_path / "sample.npz" + output_dir = tmp_path / "weathernext" + np.savez( + input_path, + input_state=np.zeros((1, 3, 4, 2), dtype=np.float32), + forcings=np.zeros((1, 3, 4, 1), dtype=np.float32), + sample_noise=np.zeros((1, 3, 4, 1), dtype=np.float32), + ) + + env = os.environ.copy() + env["PYTHONPATH"] = str(repo_root / "src") + subprocess.run( + [ + sys.executable, + str(repo_root / "examples" / "weathernext.py"), + str(output_dir), + "--input-data", + str(input_path), + "--mesh-nodes", + "5", + "--hidden-size", + "8", + ], + check=True, + cwd=repo_root, + env=env, + ) + + model = ir.load(output_dir / "model.onnx") + assert [value.name for value in model.graph.inputs] == [ + "input_state", + "forcings", + "sample_noise", + ] + assert [value.name for value in model.graph.outputs] == ["next_state"] + assert model.graph.outputs[0].shape == ir.Shape(["batch", 3, 4, 2]) + + +def test_weathernext_example_runs_f16_real_npz_inputs(tmp_path): + pytest.importorskip("onnxruntime") + + repo_root = Path(__file__).parents[1] + input_path = tmp_path / "sample.npz" + output_dir = tmp_path / "weathernext-f16" + np.savez( + input_path, + input_state=np.zeros((1, 3, 4, 2), dtype=np.float32), + forcings=np.zeros((1, 3, 4, 1), dtype=np.float32), + sample_noise=np.zeros((1, 3, 4, 1), dtype=np.float32), + ) + + env = os.environ.copy() + env["PYTHONPATH"] = str(repo_root / "src") + result = subprocess.run( + [ + sys.executable, + str(repo_root / "examples" / "weathernext.py"), + str(output_dir), + "--input-data", + str(input_path), + "--mesh-nodes", + "5", + "--hidden-size", + "8", + "--dtype", + "f16", + "--run", + "--validate", + ], + check=True, + cwd=repo_root, + env=env, + text=True, + capture_output=True, + ) + + assert "Inference output next_state shape: (1, 3, 4, 2)" in result.stdout