Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 16 additions & 2 deletions aurora/model/aurora.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,7 +26,7 @@
from aurora.model.encoder import Perceiver3DEncoder
from aurora.model.lora import LoRAMode
from aurora.model.perceiver import PerceiverAttention
from aurora.model.swin3d import Swin3DTransformerBackbone, WindowAttention
from aurora.model.swin3d import NoiseGenerator, Swin3DTransformerBackbone, WindowAttention
from aurora.normalisation import log_transform, log_untransform

__all__ = [
Expand Down Expand Up @@ -339,14 +339,27 @@ def set_noise_accumulation(self, n: int = 0) -> None:
"""
self.backbone.set_noise_accumulation(n)

def forward(self, batch: Batch, lead_times: Optional[torch.Tensor] = None) -> Batch:
def forward(
self,
batch: Batch,
lead_times: Optional[torch.Tensor] = None,
*,
generator: NoiseGenerator = None,
) -> Batch:
"""Forward pass.

Args:
batch (:class:`aurora.Batch`): Batch to run the model on.
lead_times (:class:`torch.Tensor`, optional): Per-sample lead times of shape
`(batch,)` in hours. Required when the model was configured with
`variable_lead_time=True`. Ignored otherwise.
generator (:class:`torch.Generator` or tuple of :class:`torch.Generator` or `None`,
optional): Generator for the noise in stochastic mode. A single generator is used
for the whole batch. A tuple gives one generator per batch element, so that the
noise of an element does not depend on the other elements in the batch; entries can
be `None` to use the global RNG. Generators must be on the device of the model. To
reproduce a run, re-seed the generators and call :meth:`reset_noise`. Ignored when
the model is not stochastic. Defaults to `None`, which uses the global RNG.

Returns:
:class:`Batch`: Prediction for the batch.
Expand Down Expand Up @@ -433,6 +446,7 @@ def forward(self, batch: Batch, lead_times: Optional[torch.Tensor] = None) -> Ba
lead_times=lead_times,
patch_res=patch_res,
rollout_step=batch.metadata.rollout_step,
generator=generator,
)
with context_decoder:
pred = self.decoder(
Expand Down
32 changes: 28 additions & 4 deletions aurora/model/swin3d.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,7 @@
import itertools
import warnings
from functools import lru_cache
from typing import Optional
from typing import Optional, TypeAlias

import torch
import torch.nn as nn
Expand All @@ -26,7 +26,9 @@
maybe_adjust_windows,
)

__all__ = ["Swin3DTransformerBackbone"]
__all__ = ["Swin3DTransformerBackbone", "NoiseGenerator"]

NoiseGenerator: TypeAlias = torch.Generator | tuple[torch.Generator | None, ...] | None


class MLP(nn.Module):
Expand Down Expand Up @@ -946,6 +948,25 @@ def set_noise_accumulation(self, n: int = 0) -> None:
self._noise_cache_size = max(n, 0)
self._accumulate_noise = self._noise_cache_size > 0

def _sample_noise(
self,
shape: tuple[int, ...],
device: torch.device,
dtype: torch.dtype,
generator: NoiseGenerator = None,
) -> torch.Tensor:
"""Draw noise of shape `shape`, one draw per batch element if `generator` is a tuple."""
if isinstance(generator, tuple):
if len(generator) != shape[0]:
raise ValueError(
f"Expected one generator per batch element, but got `{len(generator)}` "
f"generators for a batch of size `{shape[0]}`."
)
return torch.stack(
[torch.randn(shape[1:], device=device, dtype=dtype, generator=g) for g in generator]
)
Comment thread
wesselb marked this conversation as resolved.
return torch.randn(shape, device=device, dtype=dtype, generator=generator)

def get_encoder_specs(
self, patch_res: tuple[int, int, int]
) -> tuple[list[tuple[int, int, int]], list[tuple[int, int, int]]]:
Expand All @@ -968,6 +989,7 @@ def forward(
lead_times: torch.Tensor,
rollout_step: int,
patch_res: tuple[int, int, int],
generator: NoiseGenerator = None,
) -> torch.Tensor:
"""Run the backbone.

Expand All @@ -976,6 +998,8 @@ def forward(
lead_times (torch.Tensor): Lead times of shape `(batch,)` in hours.
rollout_step (int): Roll-out step.
patch_res (tuple[int, int, int]): Patch resolution of the form `(C, H, W)`.
generator (torch.Generator or tuple[torch.Generator | None, ...], optional): Generator
for the noise in stochastic mode. See :meth:`Aurora.forward`. Defaults to `None`.

Returns:
torch.Tensor: Output tokens of shape `(B, L, D)`.
Expand All @@ -997,7 +1021,7 @@ def forward(

if self.stochastic:
noise_shape = x.shape[:-1] + (self.embed_dim,)
noise = torch.randn(noise_shape, device=x.device, dtype=x.dtype)
noise = self._sample_noise(noise_shape, x.device, x.dtype, generator)
if self._accumulate_noise:
# Shape change (e.g. different batch size) invalidates the cache.
if self._noise_cache and self._noise_cache[0].shape != noise.shape:
Expand All @@ -1014,7 +1038,7 @@ def forward(
# Fill any remaining slots so the cache is always exactly N entries.
while len(self._noise_cache) < self._noise_cache_size:
self._noise_cache.append(
torch.randn(noise_shape, device=x.device, dtype=x.dtype)
self._sample_noise(noise_shape, x.device, x.dtype, generator)
)
effective_noise = torch.stack(self._noise_cache).sum(dim=0) / (
self._noise_cache_size**0.5
Expand Down
9 changes: 7 additions & 2 deletions aurora/rollout.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@

from aurora.batch import Batch
from aurora.model.aurora import Aurora
from aurora.model.swin3d import NoiseGenerator

__all__ = ["rollout"]

Expand Down Expand Up @@ -48,6 +49,7 @@ def rollout(
fine_lead_times: Optional[Sequence[float]] = None,
use_noise_accumulation: bool = True,
apply_rollout_input_clipping: bool = True,
generator: NoiseGenerator = None,
) -> Generator[Batch, None, None]:
"""Perform a roll-out to make long-term predictions.

Expand Down Expand Up @@ -90,6 +92,9 @@ def rollout(
back into the model during roll-out, but may be undesirable if the model was not trained
with clipping and the user wants to preserve the raw model predictions for analysis.
Default: `True`.
generator (:class:`torch.Generator` or tuple of :class:`torch.Generator` or `None`,
optional): Generator for the noise in stochastic mode, passed to every forward pass.
See :meth:`aurora.Aurora.forward`. Default: `None`.

Yields:
:class:`aurora.Batch`: The prediction after every (sub-)step.
Expand Down Expand Up @@ -128,7 +133,7 @@ def rollout(
# Inner loop: iterate over sub-step lead times.
for lt_hours in fine_lead_times:
sub_lead_times = _make_lead_time_tensor(batch, lt_hours)
pred = model.forward(batch, lead_times=sub_lead_times)
pred = model.forward(batch, lead_times=sub_lead_times, generator=generator)

yield pred

Expand All @@ -137,7 +142,7 @@ def rollout(
pred = model.apply_rollout_input_clipping(pred)
batch = _advance_batch(batch, pred)
else:
pred = model.forward(batch, lead_times=base_lead_times)
pred = model.forward(batch, lead_times=base_lead_times, generator=generator)

yield pred

Expand Down
29 changes: 29 additions & 0 deletions docs/models.md
Original file line number Diff line number Diff line change
Expand Up @@ -450,3 +450,32 @@ When using `rollout` with `fine_lead_times`, noise accumulation is enabled by de
smoother intra-step transitions while using independent effective noise between main steps,
matching the training regimen. Set `use_noise_accumulation=False` to draw independent
noise at each sub-step instead, though this is not recommended.

### Reproducible Noise

By default, the noise is drawn from the global RNG. To control the noise, pass a `torch.Generator`
on the device of the model to `Aurora.forward` or `rollout`:

```python
device = next(model.parameters()).device
generator = torch.Generator(device=device).manual_seed(42)

with torch.inference_mode():
preds = [pred.to("cpu") for pred in rollout(model, batch, steps=4, generator=generator)]
```

When generating multiple ensemble members simultaneously by using a batch size, pass a tuple with
one generator per batch element to control the noise of every member separately. The noise of a
member then does not depend on the other members in the batch. Entries can be `None` to use the
global RNG for that member.

```python
# `batch` contains three ensemble members.
generators = tuple(torch.Generator(device=device).manual_seed(seed) for seed in (1, 2, 3))

with torch.inference_mode():
preds = [pred.to("cpu") for pred in rollout(model, batch, steps=4, generator=generators)]
```

To reproduce a run, re-seed the generators and call `model.reset_noise()`, which clears noise
cached by noise accumulation.
7 changes: 4 additions & 3 deletions tests/v1p5/_helpers.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,16 +26,17 @@ def _make_batch(
surf_vars: tuple[str, ...] = _SURF_VARS,
static_vars: tuple[str, ...] = _STATIC_VARS,
atmos_vars: tuple[str, ...] = _ATMOS_VARS,
batch_size: int = BATCH,
) -> Batch:
"""Create a minimal synthetic batch."""
return Batch(
surf_vars={k: torch.randn(BATCH, HISTORY, H, W) for k in surf_vars},
surf_vars={k: torch.randn(batch_size, HISTORY, H, W) for k in surf_vars},
static_vars={k: torch.randn(H, W) for k in static_vars},
atmos_vars={k: torch.randn(BATCH, HISTORY, N_LEVELS, H, W) for k in atmos_vars},
atmos_vars={k: torch.randn(batch_size, HISTORY, N_LEVELS, H, W) for k in atmos_vars},
metadata=Metadata(
lat=torch.linspace(90, -90, H),
lon=torch.linspace(0, 360, W + 1)[:-1],
time=(datetime(2023, 6, 15, 12, 0),),
time=(datetime(2023, 6, 15, 12, 0),) * batch_size,
atmos_levels=(100, 250, 500, 850),
),
)
Expand Down
114 changes: 114 additions & 0 deletions tests/v1p5/test_forward_generator.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,114 @@
"""Copyright (c) Microsoft Corporation. Licensed under the MIT license.

Tests for the `generator` argument of `Aurora.forward` and `rollout`.
"""

from typing import Any, Callable, Sequence
from unittest.mock import patch

import pytest
import torch

from ._helpers import _OUTPUT_ONLY_SURF, _SURF_VARS, _make_batch, _make_small_v1p5
from aurora import AuroraV1p5, rollout
from aurora.model.swin3d import NoiseGenerator

_INPUT_SURF_VARS = tuple(v for v in _SURF_VARS if v not in _OUTPUT_ONLY_SURF)


@pytest.fixture
def model() -> AuroraV1p5:
return _make_small_v1p5(stochastic=True).eval()


def _record_noise(model: AuroraV1p5, run: Callable[[], Any]) -> list[torch.Tensor]:
"""Call `run` and return the noise samples drawn."""
recorded: list[torch.Tensor] = []
sample_noise = model.backbone._sample_noise

def record(*args: Any) -> torch.Tensor:
recorded.append(sample_noise(*args))
return recorded[-1]

with torch.inference_mode(), patch.object(model.backbone, "_sample_noise", record):
run()
return recorded


def _forward_noise(
model: AuroraV1p5, generator: NoiseGenerator, batch_size: int = 1, n_forwards: int = 3
) -> list[torch.Tensor]:
"""Return the noise samples drawn by `n_forwards` forward passes."""
batch = _make_batch(surf_vars=_INPUT_SURF_VARS, batch_size=batch_size)
lead_times = torch.full((batch_size,), 6.0)
return _record_noise(
model,
lambda: [
model.forward(batch, lead_times=lead_times, generator=generator)
for _ in range(n_forwards)
],
)


def _equal_sequences_of_tensors(a: Sequence[torch.Tensor], b: Sequence[torch.Tensor]) -> bool:
return len(a) == len(b) and all(torch.equal(x, y) for x, y in zip(a, b))


def _member(recorded: Sequence[torch.Tensor], i: int) -> list[torch.Tensor]:
return [noise[i] for noise in recorded]


@pytest.mark.parametrize("n", [0, 2])
def test_single_generator(model: AuroraV1p5, n: int) -> None:
model.set_noise_accumulation(n)
generator = torch.Generator().manual_seed(42)
torch.manual_seed(0)
first = _forward_noise(model, generator)
# Re-seeding the generator and flushing the cache reproduces the noise, regardless of the
# global RNG.
generator.manual_seed(42)
model.reset_noise()
torch.manual_seed(1)
assert _equal_sequences_of_tensors(first, _forward_noise(model, generator))
# Without re-seeding, the generator keeps advancing.
model.reset_noise()
assert not _equal_sequences_of_tensors(first, _forward_noise(model, generator))


def test_tuple_of_generators(model: AuroraV1p5) -> None:
def forward_noise(seeds: Sequence[int | None], batch_size: int) -> list[torch.Tensor]:
torch.manual_seed(0)
generators = tuple(None if s is None else torch.Generator().manual_seed(s) for s in seeds)
return _forward_noise(model, generators, batch_size=batch_size)

first = forward_noise((1, None, 3), batch_size=3)
second = forward_noise((2, None, 3), batch_size=3)
alone = forward_noise((3,), batch_size=1)
# The noise of a member depends only on its own generator, ...
assert not _equal_sequences_of_tensors(_member(first, 0), _member(second, 0))
assert _equal_sequences_of_tensors(_member(first, 2), _member(second, 2))
assert _equal_sequences_of_tensors(_member(first, 2), _member(alone, 0))
# ... and a `None` entry uses the global RNG.
assert _equal_sequences_of_tensors(_member(first, 1), _member(second, 1))


def test_tuple_length_mismatch(model: AuroraV1p5) -> None:
batch = _make_batch(surf_vars=_INPUT_SURF_VARS, batch_size=2)
with torch.inference_mode(), pytest.raises(ValueError, match="one generator per batch element"):
model.forward(batch, lead_times=torch.full((2,), 6.0), generator=(torch.Generator(),))


@pytest.mark.parametrize("fine_lead_times", [None, [3.0, 6.0]])
def test_rollout(model: AuroraV1p5, fine_lead_times: list[float] | None) -> None:
generator = torch.Generator().manual_seed(42)
batch = _make_batch(surf_vars=_INPUT_SURF_VARS)

def run() -> None:
list(rollout(model, batch, steps=2, fine_lead_times=fine_lead_times, generator=generator))

first = _record_noise(model, run)
# Re-seeding the generator reproduces the noise, ...
generator.manual_seed(42)
assert _equal_sequences_of_tensors(first, _record_noise(model, run))
# ... and without re-seeding, the generator keeps advancing.
assert not _equal_sequences_of_tensors(first, _record_noise(model, run))
Loading