Supported Input and Output Variables in WeatherNext's Data Modalities

WeatherNext's data_modalities module defines a flexible hierarchy of data containers that standardize five core input variables (data, lat, lon, mask, metadata) and seven output properties (point_dims_shape, data, lat, lon, mask, masked_data, metadata) across all spatial and non-spatial data types.

The WeatherNext AI weather forecasting system uses a modality-based architecture to handle diverse spatial representations—from global scalar values to irregular triangular meshes. All modalities in weathernext/utils/data_modalities.py inherit from a common abstract base class that enforces consistent input and output contracts, enabling interchangeable data pipelines regardless of underlying geometry.

Core Input Variables in Data Modalities

Every concrete modality accepts a standardized set of input variables during construction. These define the raw tensors and metadata that feed into the model.

The Five Required Input Fields

Field Type Description Always Required?
data ArrayTree Arbitrary tree of tensors (temperature, humidity, wind, etc.) Yes
lat Array Latitude coordinates in [-90, 90] Spatial modalities only
lon Array Longitude coordinates in [0, 360) Spatial modalities only
mask bool Boolean mask indicating valid points Optional (defaults to all-valid)
metadata dict Free-form auxiliary information Optional

The base class Data[ArrayTree] (lines 79-84 of [data_modalities.py](https://github.com/google-deepmind/weathernext/blob/main/weathernext/utils/data_modalities.py)) mandates implementations of point_dims_shape, data, mask, and metadata in all subclasses.

Combined Input Structure: CombinedArrays

Model internals use CombinedArrays (lines 35-43) to pack the main tensor with conditioning information:

from weathernext.utils import data_modalities as dm

# CombinedArrays is used by linear and normalization layers

combined = dm.CombinedArrays(
    data=main_tensor,                    # primary input

    norm_conditioning=norm_tensor,       # optional: norm conditioning

    other_conditioning=other_tensor      # optional: other conditioning

)

This structure allows WeatherNext models to incorporate external conditioning signals alongside core atmospheric variables.

Core Output Variables (Read-Only Properties)

The Data base class exposes output variables as read-only properties. Downstream components—loss functions, diagnostic tools, and inference pipelines—consume these consistently across all modality types.

Base Output Properties

  • point_dims_shape — Shape of leading "point" dimensions (e.g., (batch,), (batch, lat, lon), (num_points, batch))
  • data — The raw data tree (identical to input data)
  • lat / lon — Coordinate arrays (spatial modalities)
  • mask — Boolean validity mask
  • masked_data — Data with invalid points zero-filled via mask_invalid_points
  • metadata — Stored auxiliary dictionary

Modality-Specific Output Properties

Specialized modalities add convenience properties:

  • LatLonPointsData.major_axis — Axis used for flattening (line 94-95)
  • TriangularMeshData.finest_faces — Faces at finest mesh resolution

Input/Output Variable Examples by Modality Type

Global Data (Non-Spatial)

GlobalData handles scalar quantities with no spatial coordinates.

import jax.numpy as jnp
from weathernext.utils import data_modalities as dm

# Input: data (required), metadata (optional)

global_data = dm.GlobalData(
    data=jnp.ones((8,))               # shape: (batch,)

)

# Output properties

print(global_data.point_dims_shape)   # → (8,)

print(global_data.mask)               # → all-True (default)

print(global_data.masked_data.shape)  # → (8,)

Input variables used: data only
Output variables available: point_dims_shape, data, mask, masked_data, metadata

Latitude-Longitude Grid Data

LatLonGridData represents regular grids with (batch, lat, lon) structure.

import numpy as np

# Inputs: data, lat, lon, plus optional mask/metadata

lat = np.linspace(-90, 90, 32).reshape(1, 32, 1)
lon = np.linspace(0, 360, 64).reshape(1, 1, 64)

grid_data = dm.LatLonGridData(
    data=jnp.arange(8 * 32 * 64).reshape(8, 32, 64),
    lat=lat,
    lon=lon,
)

# Coordinate validation enforced: lat ∈ [-90, 90], lon ∈ [0, 360)

# Output: spatial point dimensions preserved

print(grid_data.point_dims_shape)     # → (8, 32, 64)

The helper _verify_lat_lon (lines 390-406) validates coordinate ranges across all spatial modalities.

Flattened Point Data

LatLonPointsData compresses spatial dimensions for point-based architectures.


# Input: data with shape (num_points, batch), plus coordinates

points = dm.LatLonPointsData(
    data=jnp.ones((128, 8)),          # 128 points, batch-size 8

    lat=jnp.linspace(-90, 90, 128)[:, None],
    lon=jnp.linspace(0, 360, 128)[:, None],
)

# Outputs include modality-specific convenience property

print(points.point_dims_shape)        # → (128, 8)

print(points.major_axis)              # axis used for flattening

Triangular Mesh Data

TriangularMeshData supports unstructured meshes like icosahedral discretizations.


# Build mesh with three refinement levels

tri_mesh = dm.TriangularMeshData.with_icosahedral_mesh(
    splits_list=[3, 4, 5],
    data=jnp.ones((10242, 8)),       # vertices × batch

    lat=jnp.linspace(-90, 90, 10242)[:, None],
    lon=jnp.linspace(0, 360, 10242)[:, None],
)

# Modality-specific output: finest-level faces

print(tri_mesh.finest_faces.shape)    # mesh topology at highest resolution

Mask Handling: Dense vs. Per-Axis

Masks can be supplied as dense boolean arrays or per-axis masks for memory efficiency. The internal methods _initialize_mask (lines 1088-1093) and _simplify_mask normalize these to a full-size boolean array:


# Dense mask (explicit per-point validity)

dense_mask = jnp.ones((8, 32, 64), dtype=bool)

# Per-axis mask (broadcasted automatically)

per_axis_mask = {
    1: jnp.ones(32, dtype=bool),      # lat axis

    2: jnp.ones(64, dtype=bool)       # lon axis

}

Both forms produce identical mask and masked_data outputs.

Configuration Integration

WeatherNext model configs define which modality fields correspond to physical variables. In [WeatherNext2.json](https://github.com/google-deepmind/weathernext/blob/main/weathernext/weathernext2/configs/WeatherNext2.json), the input_variables and output_variables lists map to the data field of the appropriate modality.

The model implementation in [fgn.py](https://github.com/google-deepmind/weathernext/blob/main/weathernext/weathernext2/fgn.py) accesses these through standard property names:


# Model code consumes modalities through consistent API

inputs.data      # ArrayTree of atmospheric variables

inputs.mask      # validity for loss masking

inputs.metadata  # auxiliary information

Summary

  • Five input variables (data, lat, lon, mask, metadata) standardize construction across all WeatherNext modalities
  • Seven base output variables (point_dims_shape, data, lat, lon, mask, masked_data, metadata) provide consistent downstream access
  • Modality-specific outputs (e.g., major_axis, finest_faces) extend the base contract for specialized geometries
  • Mask flexibility supports both dense and per-axis specifications with automatic normalization
  • Coordinate validation enforces lat ∈ [-90, 90] and lon ∈ [0, 360) for all spatial types

Frequently Asked Questions

What is the difference between data and masked_data outputs?

The data output returns the raw input tensor unchanged, while masked_data returns a copy with invalid points (where mask is False) set to zero via mask_invalid_points. Use masked_data for safe arithmetic operations without NaN propagation.

Can I use GlobalData for spatial quantities?

No—GlobalData lacks lat and lon fields and cannot represent spatial variation. For scalar global values like total atmospheric mass or integrated energy, use GlobalData. For any latitude/longitude-dependent field, use LatLonGridData, LatLonPointsData, or TriangularMeshData.

How does point_dims_shape differ across modality types?

point_dims_shape reflects the leading dimensions that index "points" in the data: (batch,) for GlobalData, (batch, lat, lon) for LatLonGridData, and (num_points, batch) for LatLonPointsData. This property lets downstream code reason about spatial structure without inspecting concrete class types.

Have a question about this repo?

These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →