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_fnsmust exactly match the variable names in the model's outputxr.Dataset. - Partial function storage: The constructor uses
functools.partialto 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_fnsinForwardPass.__init__enables variable-specific output transformations.- Each entry maps a variable name to
(activation_function, kwargs); the constructor binds these as partials atarchitecture.pylines 143–144. - Apply standard JAX activations directly, or use
shifted_activationfromweathernext.utils.activationsfor 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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →