# How to Customize PointsMeshUpdateConstructor for Different Data Modalities in WeatherNext

> Customize PointsMeshUpdateConstructor for diverse data modalities in WeatherNext. Inherit SpatialData, configure connectivity, and tune hyperparameters for optimal results.

- Repository: [Google DeepMind/weathernext](https://github.com/google-deepmind/weathernext)
- Tags: how-to-guide
- Published: 2026-08-12

---

**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`](https://github.com/google-deepmind/weathernext/blob/main/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`](https://github.com/google-deepmind/weathernext/blob/main/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`](https://github.com/google-deepmind/weathernext/blob/main/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`](https://github.com/google-deepmind/weathernext/blob/main/data_modalities.py):

```python
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`](https://github.com/google-deepmind/weathernext/blob/main/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:

```python
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:

```python

# 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`**

```python

# 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`**

```python
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`](https://github.com/google-deepmind/weathernext/blob/main/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`](https://github.com/google-deepmind/weathernext/blob/main/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`](https://github.com/google-deepmind/weathernext/blob/main/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`](https://github.com/google-deepmind/weathernext/blob/main/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.