diff --git a/aurora/model/aurora.py b/aurora/model/aurora.py index a668f65c..4db2a09d 100644 --- a/aurora/model/aurora.py +++ b/aurora/model/aurora.py @@ -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__ = [ @@ -339,7 +339,13 @@ 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: @@ -347,6 +353,13 @@ def forward(self, batch: Batch, lead_times: Optional[torch.Tensor] = None) -> Ba 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. @@ -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( diff --git a/aurora/model/swin3d.py b/aurora/model/swin3d.py index 58ded1e9..db55816a 100644 --- a/aurora/model/swin3d.py +++ b/aurora/model/swin3d.py @@ -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 @@ -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): @@ -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] + ) + 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]]]: @@ -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. @@ -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)`. @@ -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: @@ -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 diff --git a/aurora/rollout.py b/aurora/rollout.py index 3a84ce6b..bcd133d4 100644 --- a/aurora/rollout.py +++ b/aurora/rollout.py @@ -8,6 +8,7 @@ from aurora.batch import Batch from aurora.model.aurora import Aurora +from aurora.model.swin3d import NoiseGenerator __all__ = ["rollout"] @@ -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. @@ -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. @@ -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 @@ -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 diff --git a/docs/models.md b/docs/models.md index 3b4bc8d1..e0613717 100644 --- a/docs/models.md +++ b/docs/models.md @@ -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. diff --git a/tests/v1p5/_helpers.py b/tests/v1p5/_helpers.py index 4ddd7a65..575f357b 100644 --- a/tests/v1p5/_helpers.py +++ b/tests/v1p5/_helpers.py @@ -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), ), ) diff --git a/tests/v1p5/test_forward_generator.py b/tests/v1p5/test_forward_generator.py new file mode 100644 index 00000000..288ed094 --- /dev/null +++ b/tests/v1p5/test_forward_generator.py @@ -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))