How to Configure Per-Variable Activation Functions in WeatherNext's Output Layer

Pass a per_var_activation_fns dictionary when constructing the ForwardPass module to apply distinct activation functions to individual forecast variables.

WeatherNext supports per-variable activation functions in its output layer, allowing you to apply tailored transformations to each predicted variable—such as sigmoid for precipitation or relu for temperature. This guide explains how to configure these activations using the per_var_activation_fns parameter in weathernext2.architecture.ForwardPass, according to the google-deepmind/weathernext source code.


Understanding the per_var_activation_fns Parameter

The ForwardPass constructor accepts per_var_activation_fns as a mapping from variable names to tuples of (activation_function, activation_kwargs).

In weathernext/weathernext2/architecture.py at line 67, the parameter is defined as:

per_var_activation_fns: Optional[Mapping[str, Tuple[Callable[..., Any], Mapping[str, Any]]]] = None

The constructor stores these as partial functions that bind the provided keyword arguments to each activation (lines 143–144). During the forward pass, WeatherNext iterates over this dictionary and applies the stored callables to the corresponding variables (lines 285–287).

Each activation is applied after decoding the model's raw predictions but before returning the final xr.Dataset.


Supported Activation Functions

WeatherNext provides activation utilities in weathernext/utils/activations.py:

Function Purpose Location
shifted_activation Generic wrapper for JAX activations with optional input_offset and output_scale lines 23–39
sigmoid Ready-to-use sigmoid activation lines 42–45

You can also use any JAX-native activation: jax.nn.relu, jax.nn.tanh, jax.nn.softplus, etc.


Configuration Methods

Method 1: Direct Instantiation in Python

Build the per_var_activation_fns dictionary with variable names as keys and (function, kwargs) tuples as values.

Example: Sigmoid for Precipitation, ReLU for Temperature

import jax
import jax.numpy as jnp
from weathernext.weathernext2.architecture import ForwardPass
from weathernext.utils import activations as act

per_var_activations = {
    "precipitation": (act.sigmoid, {}),          # sigmoid with default args

    "temperature": (jax.nn.relu, {}),            # raw ReLU, no extra args

}

model = ForwardPass(
    latent_dense_kwargs=latent_kwargs,
    output_dense_kwargs=output_kwargs,
    spatial_features_kwargs=spatial_kwargs,
    mesh_num_splits=3,
    points_to_mesh_model_ctor=points_to_mesh_ctor,
    mesh_model_ctor=mesh_ctor,
    mesh_to_grid_model_ctor=mesh_to_grid_ctor,
    per_var_activation_fns=per_var_activations,
)

Method 2: Using shifted_activation for Custom Scaling

The shifted_activation helper lets you offset inputs and scale outputs—useful for variables with specific physical constraints.

from weathernext.utils.activations import shifted_activation

per_var_activations = {
    "humidity": (
        shifted_activation,
        dict(
            activation_fn=jax.nn.tanh,
            input_offset=0.1,      # shift input upward before activation

            output_scale=2.0,      # amplify output after activation

        ),
    ),
}

model = ForwardPass(
    ...,
    per_var_activation_fns=per_var_activations,
)

Method 3: Configuration via fiddle Configs

When using Fiddle for hyperparameter management, embed the activation map directly in the fdl.Config:

import fiddle as fdl
import jax
from weathernext.weathernext2 import architecture
from weathernext.utils import activations

cfg = fdl.Config(
    architecture.ForwardPass,
    latent_dense_kwargs=latent_kwargs,
    output_dense_kwargs=output_kwargs,
    spatial_features_kwargs=spatial_kwargs,
    mesh_num_splits=3,
    points_to_mesh_model_ctor=points_to_mesh_ctor,
    mesh_model_ctor=mesh_ctor,
    mesh_to_grid_model_ctor=mesh_to_grid_ctor,
    per_var_activation_fns={
        "wind_speed": (jax.nn.softplus, {}),
        "solar_radiation": (
            activations.shifted_activation,
            {
                "activation_fn": jax.nn.relu,
                "output_scale": 0.5,
            },
        ),
    },
)

Key Implementation Details

  • Variable name matching: Keys in per_var_activation_fns must exactly match the variable names in the model's output xr.Dataset.
  • Partial function storage: The constructor uses functools.partial to freeze kwargs, eliminating per-step overhead.
  • Activation timing: Transformations occur after the final decoder but before dataset packaging—ensuring gradients flow through learned parameters.

Summary

  • per_var_activation_fns in ForwardPass.__init__ enables variable-specific output transformations.
  • Each entry maps a variable name to (activation_function, kwargs); the constructor binds these as partials at architecture.py lines 143–144.
  • Apply standard JAX activations directly, or use shifted_activation from weathernext.utils.activations for offset/scaling control.
  • Configure via direct Python instantiation or Fiddle configs for experiment management.

Frequently Asked Questions

What happens if a variable name doesn't match any key in per_var_activation_fns?

Variables without an entry in the mapping receive no activation—they pass through unchanged. Only explicitly listed variables are transformed during the forward pass at architecture.py lines 285–287.

Can I use custom activation functions not in weathernext.utils.activations?

Yes. Any callable compatible with JAX's transformation rules works. The function signature should accept jax.Array and optional kwargs. Pass your custom function exactly as you would a built-in JAX activation.

How do per-variable activations affect model training?

They participate fully in gradient computation since they're applied within the ForwardPass forward method. The partial-bound activations are differentiable through both the activation itself and any shifted_activation parameters like input_offset.

When should I use sigmoid versus shifted_activation with output_scale?

Use sigmoid when you need outputs strictly bounded in (0, 1). Use shifted_activation with output_scale when you need bounded outputs with a custom range—for example, scaling sigmoid or tanh outputs to match physical units like mm/hr precipitation.

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 →