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

> Learn how to configure per-variable activation functions in WeatherNext's output layer. Apply distinct functions to individual forecast variables using the per_var_activation_fns dictionary.

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

---

**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`](https://github.com/google-deepmind/weathernext/blob/main/weathernext/weathernext2/architecture.py) at line 67, the parameter is defined as:

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

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

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

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