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 in data_modalities.py) to define new point or mesh modalities with required lat, lon, mask, data, and replace_data interfaces
  • Pass modality identifiers via points_name and mesh_name to PointsMeshTypedGraphGNN (lines 81-101 in points_mesh_gnn.py)
  • Select direction with is_points_to_mesh/is_mesh_to_points flags 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_kwargs and deep_gnn_kwargs to your feature dimensionality and complexity requirements
  • Subclass _get_edge_indices for 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:

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 →