How to Customize PointsMeshUpdateConstructor for Different Data Modalities in WeatherNext
To customize PointsMeshUpdateConstructor for different data modalities in WeatherNext, inherit from SpatialData to create new modality classes, then pass the modality identifiers and appropriate connectivity settings to PointsMeshTypedGraphGNN while adjusting layer hyperparameters to match your data's dimensionality.
The PointsMeshTypedGraphGNN class in weathernext/utils/points_mesh_gnn.py implements the core mechanism for passing messages between point-based observations and mesh-based representations. Understanding how to extend this constructor is essential when working with custom sensor networks, alternative mesh discretizations, or non-geographic data types within the WeatherNext framework.
Understanding the PointsMeshUpdateConstructor Architecture
The constructor builds a bipartite TypedGraph GNN that connects two modalities: a set of points (such as latitude-longitude observations) and a triangular mesh (such as an icosahedral discretization). This architecture enables flexible message passing in either direction through shared edge-encoding and deep message-passing components.
Core Constructor Arguments
According to the WeatherNext source code in utils/points_mesh_gnn.py (lines 81-101), the constructor accepts:
| Argument | Purpose | Typical Value |
|---|---|---|
name |
Base name with {points_name} and {mesh_name} placeholders |
"points_{points_name}_to_mesh_{mesh_name}" |
points_name |
Identifier for the points modality | "satellite" or custom |
mesh_name |
Identifier for the mesh modality | "icosahedral" or custom |
is_points_to_mesh / is_mesh_to_points |
Direction flags (exactly one must be True) |
True/False |
connectivity_type |
Edge construction strategy: CLOSEST, IN_TRIANGLE, or BALL_QUERY |
ConnectivityType.CLOSEST |
dense_kwargs |
Dense layer configuration for the underlying DeepGNN |
{"output_size": 64, "activation": "relu"} |
deep_gnn_kwargs |
Arguments forwarded to DeepGNN |
{"num_message_passing_steps": 3} |
spatial_edge_features_kwargs |
Parameters for spatial edge-feature encoding | {"use_distance": True} |
stacked_points_inputs |
Whether points use stacked ensemble format [points, batch, stack, features] |
False |
edge_encoder_dense_kwargs |
Optional dense kwargs for edge encoder | Defaults to dense_kwargs |
ball_query_radius_fraction |
Relative radius for BALL_QUERY connectivity |
0.1 |
The constructor validates these arguments—for example, raising an error if ball_query_radius_fraction is missing when using BALL_QUERY connectivity—and instantiates two key components: a DeepGNN-based typed graph and an edge encoder for spatial features.
Required Data Modality Interfaces
The two modalities that PointsMeshTypedGraphGNN expects are defined in utils/data_modalities.py:
- LatLonPointsData (lines 668-690): A flat list of points with latitude/longitude coordinates
- TriangularMeshData (lines 998-1020): Vertices forming a triangular mesh with optional face-sets
Both inherit from SpatialData, which enforces that data carries lat, lon, and a broadcastable mask. Any custom modality must implement these properties plus data, point_dims_shape, and replace_data.
Step-by-Step Customization Guide
Step 1: Create a Custom Points Modality
To add a new point-based data source, define a class inheriting from SpatialData as demonstrated in data_modalities.py:
from weathernext.utils import data_modalities
import chex
import numpy as np
class SensorStationData(data_modalities.SpatialData[chex.Array]):
"""Custom modality for sensor station observations."""
def __init__(self, data: chex.Array, lat: np.ndarray, lon: np.ndarray,
mask: np.ndarray | None = None):
all_shared = {
"data": data,
"lat": lat,
"lon": lon,
"mask": mask if mask is not None else np.ones(lat.shape, bool)
}
self._container = data_modalities.ShareLeadingAxesArrayTree(
all_shared, num_leading_shared_axes=2)
self._metadata = {}
@property
def point_dims_shape(self) -> tuple[int, int]:
return self._container.shared_leading_shape
@property
def data(self):
return self._container.array_tree["data"]
@property
def lat(self):
return self._container.array_tree["lat"]
@property
def lon(self):
return self._container.array_tree["lon"]
@property
def mask(self):
return self._container.array_tree["mask"]
@property
def metadata(self):
return self._metadata
def replace_data(self, data, *, lat=None, lon=None, **kw):
return SensorStationData(
data=data,
lat=lat if lat is not None else self.lat,
lon=lon if lon is not None else self.lon,
mask=self.mask,
)
This implementation follows the same pattern as LatLonPointsData in lines 668-690 of data_modalities.py, using ShareLeadingAxesArrayTree to handle batched data efficiently.
Step 2: Instantiate PointsMeshTypedGraphGNN with Custom Identifiers
Pass your custom modality names and configure the connectivity:
from weathernext.utils.points_mesh_gnn import PointsMeshTypedGraphGNN, ConnectivityType
import haiku as hk
# Layer configuration tailored to sensor data dimensionality
dense_cfg = {"output_size": 128, "activation": "relu"}
deep_cfg = {"num_message_passing_steps": 2, "hidden_size": 64}
edge_feat_cfg = {"use_distance": True}
# Custom PointsMeshUpdateConstructor instantiation
points_mesh_update = PointsMeshTypedGraphGNN(
name="sensor_to_mesh_{points_name}_{mesh_name}",
points_name="sensor", # matches our custom modality
mesh_name="icosahedral",
is_points_to_mesh=True, # sensor stations send to mesh
is_mesh_to_points=False,
connectivity_type=ConnectivityType.CLOSEST,
dense_kwargs=dense_cfg,
deep_gnn_kwargs=deep_cfg,
spatial_edge_features_kwargs=edge_feat_cfg,
stacked_points_inputs=False,
ball_query_radius_fraction=None,
)
# Standard Haiku transform for integration
points_mesh_fn = hk.transform(lambda mesh, points: points_mesh_update(mesh, points))
The name parameter uses placeholders that get formatted with actual modality names, enabling consistent block naming across architectures.
Step 3: Use Custom Mesh Modalities
For mesh alternatives (Voronoi tessellations, adaptive meshes), provide any object satisfying the TriangularMeshData interface:
# Construct from Voronoi generator output
voronoi_mesh = data_modalities.TriangularMeshData(
data=mesh_features,
lat=mesh_latitudes[:, None],
lon=mesh_longitudes[:, None],
face_sets=[("voronoi", voronoi_faces)], # finest mesh face set
)
# Reverse direction: mesh sends to points
mesh_to_sensor = PointsMeshTypedGraphGNN(
name="mesh_to_sensor_{points_name}_{mesh_name}",
points_name="sensor",
mesh_name="voronoi", # custom mesh identifier
is_points_to_mesh=False,
is_mesh_to_points=True,
connectivity_type=ConnectivityType.IN_TRIANGLE,
dense_kwargs=dense_cfg,
deep_gnn_kwargs=deep_cfg,
spatial_edge_features_kwargs=edge_feat_cfg,
)
The constructor only requires vertex coordinates and face lists—it does not depend on mesh generation method.
Step 4: Handle Non-Geographic Modalities
For data without natural latitude/longitude coordinates, you have two options:
Option A: Dummy coordinates with BALL_QUERY
# Use dummy coordinates and large radius for connectivity
non_geo_update = PointsMeshTypedGraphGNN(
name="features_to_mesh_{points_name}_{mesh_name}",
points_name="embedding",
mesh_name="icosahedral",
is_points_to_mesh=True,
is_mesh_to_points=False,
connectivity_type=ConnectivityType.BALL_QUERY,
ball_query_radius_fraction=10.0, # large enough to connect all
dense_kwargs=dense_cfg,
deep_gnn_kwargs=deep_cfg,
)
Option B: Subclass with custom _get_edge_indices
class FeatureSimilarityGNN(PointsMeshTypedGraphGNN):
def _get_edge_indices(self, points_latitude, points_longitude,
mesh_latitude, mesh_longitude, mesh_faces):
# Custom edge construction using feature similarity
# Access external features via closure or instance attribute
from scipy.spatial import cKDTree
tree = cKDTree(self.mesh_features) # (N_mesh, feature_dim)
_, mesh_idx = tree.query(self.point_features, k=5)
points_idx = np.repeat(
np.arange(self.point_features.shape[0]), 5)
return points_idx, mesh_idx.ravel()
This approach preserves all constructor arguments while replacing the geometric edge computation in _get_edge_indices.
Selecting Connectivity Types for Your Use Case
The ConnectivityType enum in points_mesh_gnn.py determines how edges are constructed:
- CLOSEST: Each point connects to its nearest mesh vertex. Best for sparse, irregular point distributions.
- IN_TRIANGLE: Connects points to the three vertices of containing triangles. Requires point-in-triangle testing; optimal when points are dense relative to mesh resolution.
- BALL_QUERY: Connects all vertices within a radius fraction of the finest mesh's longest edge. Most flexible for non-geographic or feature-based connectivity.
Adjusting Hyperparameters for Data Dimensionality
The dense_kwargs and deep_gnn_kwargs parameters must scale with your data:
| Data Characteristic | Recommended Adjustment |
|---|---|
| High-dimensional point features (>256) | Increase output_size in dense_kwargs |
| Complex spatial relationships | Increase num_message_passing_steps |
| Stacked ensemble inputs | Set stacked_points_inputs=True and verify batch handling |
| Large mesh/point count disparities | Tune edge_encoder_dense_kwargs separately from dense_kwargs |
Summary
- Extend
SpatialData(lines 668-690, 998-1020 indata_modalities.py) to define new point or mesh modalities with requiredlat,lon,mask,data, andreplace_datainterfaces - Pass modality identifiers via
points_nameandmesh_nametoPointsMeshTypedGraphGNN(lines 81-101 inpoints_mesh_gnn.py) - Select direction with
is_points_to_mesh/is_mesh_to_pointsflags based on which modality sends messages - Choose connectivity (
CLOSEST,IN_TRIANGLE,BALL_QUERY) matching your geometric or feature-based relationship - Scale layer parameters in
dense_kwargsanddeep_gnn_kwargsto your feature dimensionality and complexity requirements - Subclass
_get_edge_indicesfor completely custom edge construction when standard geometric strategies are insufficient
Frequently Asked Questions
What is the minimum required interface for a custom data modality?
Your class must inherit from SpatialData and implement lat, lon, mask, data, point_dims_shape, and replace_data properties/methods. The replace_data method is critical because update_blocks.update_main_data uses it to write GNN outputs back into your modality after message passing.
Can I use PointsMeshTypedGraphGNN without geographic coordinates?
Yes. Provide dummy latitude/longitude values and use ConnectivityType.BALL_QUERY with a sufficiently large radius, or subclass and override _get_edge_indices to implement feature-based or learned connectivity. The underlying TypedGraph machinery does not inherently depend on geographic meaning.
How do I handle ensemble predictions with multiple point clouds?
Set stacked_points_inputs=True in the constructor. This indicates your points data has shape [points, batch, stack, features] rather than [points, batch, features]. The internal ShareLeadingAxesArrayTree will correctly broadcast operations across the stack dimension.
Why does my custom modality fail during the update step?
Verify that replace_data returns a new instance with updated data while preserving lat, lon, and mask. The update mechanism in update_blocks.py expects immutability semantics—modifying arrays in-place will cause inconsistent state between GNN outputs and downstream processing.
What determines the memory consumption of the Points-Mesh block?
Primary factors are: (1) the product of point and mesh counts for BALL_QUERY, (2) the number of message-passing steps in deep_gnn_kwargs, and (3) the output_size dimensions. Use pad_edges_to_multiple_of only when required by distributed training sharding—otherwise leave as None to avoid unnecessary padding overhead.
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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →