From afbe29d25b9aeb5763ed3d42cfb7bcd1536762df Mon Sep 17 00:00:00 2001 From: s-sasaki-earthsea-wizard Date: Sun, 23 Aug 2026 12:42:25 +0900 Subject: [PATCH 1/6] Add generator argument to control noise in stochastic mode Support passing a torch.Generator, or a tuple with one generator per batch element, to Aurora.forward, Swin3DTransformerBackbone.forward, and rollout, so that the noise injected by stochastic models can be reproduced per ensemble member (#191). A single generator drives one stream for the whole batch, while a tuple draws each batch element from its own stream, making a member's noise sequence independent of the batch composition. The noise cache is now also invalidated on device or dtype changes, not only shape changes. --- aurora/model/aurora.py | 38 +++- aurora/model/swin3d.py | 87 ++++++++- aurora/rollout.py | 13 +- docs/models.md | 34 ++++ tests/v1p5/_helpers.py | 7 +- tests/v1p5/test_forward_generator.py | 254 +++++++++++++++++++++++++++ 6 files changed, 420 insertions(+), 13 deletions(-) create mode 100644 tests/v1p5/test_forward_generator.py diff --git a/aurora/model/aurora.py b/aurora/model/aurora.py index a668f65c..edd185e9 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__ = [ @@ -323,6 +323,9 @@ def __init__( if isinstance(m, (WindowAttention, PerceiverAttention)): m.use_fp16_safe_attention = True + # Warn only once when `generator` is passed to `forward` of a non-stochastic model. + self._generator_ignored_warned = False + def reset_noise(self) -> None: """Flush the backbone noise cache. @@ -339,7 +342,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,10 +356,34 @@ 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): Source of randomness for the noise injection in stochastic mode. A + single generator drives one stream for the whole batch. A tuple must contain one + entry per batch element (ensemble member), in the current order of the batch + dimension; every element then draws from its own stream, so the noise sequence of + a given member does not depend on the batch composition. Tuple entries may be + `None` to fall back to the global RNG for that member, and passing the same + generator object in several slots makes those members share one stream. Because + the two modes draw with different shapes, a single generator and a tuple are not + interchangeable. Generators must live on the same device as the model, and they + advance on every forward pass. To reproduce a run, re-seed the generators (e.g. + with `manual_seed`) *and* call :meth:`reset_noise`, so that noise cached by noise + accumulation in a previous run cannot contaminate the reproduced sequence. When + the model is not stochastic, this argument is ignored with a warning. Defaults to + `None`, which draws from the global RNG (the previous behaviour). Returns: :class:`Batch`: Prediction for the batch. """ + if generator is not None and not self.backbone.stochastic: + if not self._generator_ignored_warned: + warnings.warn( + "`generator` is ignored because stochastic noise is disabled.", + stacklevel=2, + ) + self._generator_ignored_warned = True + generator = None + batch = self.batch_transform_hook(batch) # Get the first parameter. We'll derive the data type and device from this parameter. @@ -433,6 +466,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..7bace233 100644 --- a/aurora/model/swin3d.py +++ b/aurora/model/swin3d.py @@ -28,6 +28,11 @@ __all__ = ["Swin3DTransformerBackbone"] +NoiseGenerator = torch.Generator | tuple[torch.Generator | None, ...] | None +"""Source of randomness for noise injection in stochastic mode: a single generator driving one +stream for the whole batch, a tuple with one generator per batch element, or `None` for the +global RNG.""" + class MLP(nn.Module): """A one-hidden-layer MLP with dropout after the hidden layer and at the end.""" @@ -922,6 +927,12 @@ def reset_noise(self) -> None: Call this to clear all cached noise tensors, e.g. at the beginning of a new forecast issue time. After reset, the next forward call starts building the cache afresh. + + To reproduce a run that controls the noise with `generator` (see :meth:`forward`), re-seed + the generators *and* call this method: cached noise left over from a previous run would + otherwise contaminate the reproduced sequence. The same applies after changing the order or + composition of the ensemble members that a tuple of generators corresponds to, which cannot + be detected from the noise tensors themselves. """ if self.stochastic: @@ -946,6 +957,57 @@ 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 _validate_generator( + self, generator: NoiseGenerator, batch_size: int, device: torch.device + ) -> None: + """Validate `generator` against the batch before any randomness is consumed. + + Checking the tuple length and every device up front ensures that no generator has already + advanced when an error is raised for a later batch element. + """ + if not isinstance(generator, tuple): + return + if len(generator) != batch_size: + raise ValueError( + f"Expected {batch_size} generators (one per batch element), got {len(generator)}." + ) + for i, g in enumerate(generator): + if g is None: + continue + # Like PyTorch, treat an index-less device (e.g. `cuda`) as compatible with an + # indexed one (e.g. `cuda:0`); only compare indices when both are explicit. + if g.device.type != device.type or ( + g.device.index is not None + and device.index is not None + and g.device.index != device.index + ): + raise ValueError( + f"Generator for batch element {i} is on device `{g.device}`, but noise is " + f"generated on device `{device}`." + ) + + def _sample_noise( + self, + shape: tuple[int, ...], + device: torch.device, + dtype: torch.dtype, + generator: NoiseGenerator = None, + ) -> torch.Tensor: + """Draw one noise sample of shape `(B, L, D)`. + + A single generator produces the whole sample in one draw, so its stream depends on the + batch size. A tuple of generators instead produces one `(L, D)` draw per batch element, so + the sequence of draws for a given element does not depend on the batch composition. `None`, + or a `None` entry in a tuple, falls back to the global RNG. Because the two modes consume + a generator with differently shaped draws, a single generator and a tuple are not + interchangeable. + """ + if isinstance(generator, tuple): + 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 +1030,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 +1039,10 @@ 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): Source of + randomness for noise injection in stochastic mode. See :meth:`_sample_noise` and + :meth:`Aurora.forward` for the semantics. Only used when the model is stochastic. + Defaults to `None`, which draws from the global RNG. Returns: torch.Tensor: Output tokens of shape `(B, L, D)`. @@ -997,13 +1064,21 @@ def forward( if self.stochastic: noise_shape = x.shape[:-1] + (self.embed_dim,) - noise = torch.randn(noise_shape, device=x.device, dtype=x.dtype) + self._validate_generator(generator, x.shape[0], x.device) + 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: + # A shape (e.g. different batch size), device, or dtype change invalidates the + # cache. + cached = self._noise_cache[0] if self._noise_cache else None + if cached is not None and ( + cached.shape != noise.shape + or cached.device != noise.device + or cached.dtype != noise.dtype + ): warnings.warn( - f"Noise shape changed from {self._noise_cache[0].shape} to " - f"{noise.shape}; clearing noise cache.", + f"Cached noise of shape {cached.shape} ({cached.dtype} on " + f"{cached.device}) is incompatible with new noise of shape {noise.shape} " + f"({noise.dtype} on {noise.device}); clearing noise cache.", stacklevel=2, ) self._noise_cache.clear() @@ -1014,7 +1089,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..ef13cc25 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,13 @@ 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): Source of randomness for the noise injection of stochastic models, passed + to every `forward` call of the roll-out. See :meth:`aurora.Aurora.forward` for the + semantics. The generators advance on every (sub-)step, and a tuple corresponds to the + order of the batch dimension throughout the roll-out. To reproduce a roll-out, re-seed + the generators and call `model.reset_noise()` before starting. Default: `None`, which + draws from the global RNG. Yields: :class:`aurora.Batch`: The prediction after every (sub-)step. @@ -128,7 +137,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 +146,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..fddb2faf 100644 --- a/docs/models.md +++ b/docs/models.md @@ -450,3 +450,37 @@ 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 injected noise is drawn from the global RNG, so the noise of an individual +ensemble member cannot easily be reproduced. To control the noise, pass a `torch.Generator` +to `Aurora.forward` or `rollout`. The generator must live on the same device as the model: + +```python +import torch + +from aurora import rollout + +device = next(model.parameters()).device +generator = torch.Generator(device=device).manual_seed(42) +preds = [pred for pred in rollout(model, batch, steps=10, generator=generator)] +``` + +When generating multiple ensemble members simultaneously by using a batch size, pass a tuple +with one generator per batch element (ensemble member) to control the noise of every member +separately. Every member then draws from its own stream, so a member's noise sequence does not +depend on the batch composition. Tuple entries may be `None` to fall back to the global RNG for +that member. Note that a single generator and a tuple of generators draw with different shapes +and are therefore not interchangeable. + +```python +# `batch` contains three ensemble members. +generators = tuple(torch.Generator(device=device).manual_seed(seed) for seed in (1, 2, 3)) +preds = [pred for pred in rollout(model, batch, steps=10, generator=generators)] +``` + +Generators advance on every forward pass. +To reproduce a run, re-seed the generators with `manual_seed` *and* call `model.reset_noise()`: +the latter clears noise cached by noise accumulation, which would otherwise contaminate the +reproduced sequence. 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..aea562dc --- /dev/null +++ b/tests/v1p5/test_forward_generator.py @@ -0,0 +1,254 @@ +"""Copyright (c) Microsoft Corporation. Licensed under the MIT license. + +Tests for the `generator` argument of `Aurora.forward`, which controls the noise injection of +stochastic Aurora 1.5 models (issue #191). +""" + +import warnings + +import pytest +import torch + +from ._helpers import _OUTPUT_ONLY_SURF, _SURF_VARS, _make_batch, _make_small_v1p5 +from aurora import rollout + +_INPUT_SURF_VARS = tuple(v for v in _SURF_VARS if v not in _OUTPUT_ONLY_SURF) + + +def _record_noise(model, n_forwards=3, generator=None, batch_size=1): + """Run `n_forwards` identical forward passes and return the noise samples drawn.""" + batch = _make_batch(surf_vars=_INPUT_SURF_VARS, batch_size=batch_size) + lead_times = torch.full((batch_size,), 6.0) + recorded = [] + original = model.backbone._sample_noise + + def recording_sample_noise(shape, device, dtype, generator=None): + noise = original(shape, device, dtype, generator) + recorded.append(noise.clone()) + return noise + + model.backbone._sample_noise = recording_sample_noise + try: + with torch.inference_mode(): + for _ in range(n_forwards): + model.forward(batch, lead_times=lead_times, generator=generator) + finally: + model.backbone._sample_noise = original + return recorded + + +def _seqs_equal(a, b): + return len(a) == len(b) and all(torch.equal(x, y) for x, y in zip(a, b)) + + +def _member_seq(recorded, i): + return [noise[i] for noise in recorded] + + +def test_single_generator_reseed_reproduces_noise(): + model = _make_small_v1p5(stochastic=True) + model.eval() + generator = torch.Generator().manual_seed(42) + first = _record_noise(model, generator=generator) + generator.manual_seed(42) + second = _record_noise(model, generator=generator) + assert _seqs_equal(first, second) + + +def test_single_generator_advances_without_reseed(): + model = _make_small_v1p5(stochastic=True) + model.eval() + generator = torch.Generator().manual_seed(42) + first = _record_noise(model, generator=generator) + second = _record_noise(model, generator=generator) + assert not _seqs_equal(first, second) + + +def test_same_seed_fresh_instance_reproduces_noise(): + # The core requirement of issue #191: a given ensemble member reproduces the same noise + # sequence across model runs, here even across fresh model instances. + model_a = _make_small_v1p5(stochastic=True) + model_b = _make_small_v1p5(stochastic=True) + model_a.eval() + model_b.eval() + generator_a = torch.Generator().manual_seed(42) + generator_b = torch.Generator().manual_seed(42) + assert _seqs_equal( + _record_noise(model_a, generator=generator_a), + _record_noise(model_b, generator=generator_b), + ) + + +def test_single_generator_independent_of_global_rng(): + model = _make_small_v1p5(stochastic=True) + model.eval() + generator = torch.Generator().manual_seed(42) + torch.manual_seed(0) + first = _record_noise(model, generator=generator) + generator.manual_seed(42) + torch.manual_seed(999) + second = _record_noise(model, generator=generator) + assert _seqs_equal(first, second) + + +def test_generator_none_preserves_global_rng_semantics(): + model = _make_small_v1p5(stochastic=True) + model.eval() + torch.manual_seed(0) + first = _record_noise(model) + torch.manual_seed(0) + second = _record_noise(model) + torch.manual_seed(999) + third = _record_noise(model) + assert _seqs_equal(first, second) + assert not _seqs_equal(first, third) + + +def test_tuple_member_noise_independent_of_batch_composition(): + model = _make_small_v1p5(stochastic=True) + model.eval() + generators = tuple(torch.Generator().manual_seed(seed) for seed in (1, 2, 3)) + joint = _record_noise(model, generator=generators, batch_size=3) + alone = _record_noise(model, generator=(torch.Generator().manual_seed(3),), batch_size=1) + assert _seqs_equal(_member_seq(joint, 2), _member_seq(alone, 0)) + + +def test_tuple_and_single_generator_are_not_interchangeable(): + model = _make_small_v1p5(stochastic=True) + model.eval() + single = _record_noise(model, generator=torch.Generator().manual_seed(42), batch_size=2) + generators = (torch.Generator().manual_seed(42), torch.Generator().manual_seed(42)) + per_member = _record_noise(model, generator=generators, batch_size=2) + # Identically seeded per-member generators give every member the same noise, whereas a single + # generator drives one stream across the whole batch. + assert _seqs_equal(_member_seq(per_member, 0), _member_seq(per_member, 1)) + assert not _seqs_equal(_member_seq(single, 0), _member_seq(single, 1)) + assert not _seqs_equal(_member_seq(single, 1), _member_seq(per_member, 1)) + + +def test_tuple_none_member_independent_of_other_generators(): + model = _make_small_v1p5(stochastic=True) + model.eval() + torch.manual_seed(123) + first = _record_noise(model, generator=(torch.Generator().manual_seed(42), None), batch_size=2) + torch.manual_seed(123) + second = _record_noise(model, generator=(torch.Generator().manual_seed(7), None), batch_size=2) + # The seeded member changes with its seed, but the `None` member draws from the global RNG + # and must not be affected by the other member's generator. + assert not _seqs_equal(_member_seq(first, 0), _member_seq(second, 0)) + assert _seqs_equal(_member_seq(first, 1), _member_seq(second, 1)) + + +def test_tuple_length_mismatch_raises_without_consuming_rng(): + model = _make_small_v1p5(stochastic=True) + model.eval() + generators = tuple(torch.Generator().manual_seed(seed) for seed in (1, 2, 3)) + states = [g.get_state().clone() for g in generators] + batch = _make_batch(surf_vars=_INPUT_SURF_VARS, batch_size=2) + with torch.inference_mode(), pytest.raises(ValueError, match="Expected 2 generators"): + model.forward(batch, lead_times=torch.full((2,), 6.0), generator=generators) + for state, g in zip(states, generators): + assert torch.equal(state, g.get_state()) + + +def test_accumulation_reproduces_after_reseed_and_reset_noise(): + model = _make_small_v1p5(stochastic=True) + model.eval() + model.set_noise_accumulation(n=2) + generator = torch.Generator().manual_seed(42) + first = _record_noise(model, generator=generator) + # Reproducing a run requires *both* re-seeding the generator and flushing the noise cache. + generator.manual_seed(42) + model.reset_noise() + second = _record_noise(model, generator=generator) + # Re-seeding alone is not enough: leftover cached noise changes how many samples are drawn. + generator.manual_seed(42) + third = _record_noise(model, generator=generator) + model.set_noise_accumulation(n=0) + assert _seqs_equal(first, second) + assert not _seqs_equal(first, third) + + +def test_non_stochastic_model_warns_once_and_ignores_generator(): + model = _make_small_v1p5() + model.eval() + generator = torch.Generator().manual_seed(42) + state = generator.get_state().clone() + batch = _make_batch(surf_vars=_INPUT_SURF_VARS) + lead_times = torch.tensor([6.0]) + with torch.inference_mode(): + with pytest.warns(UserWarning, match="`generator` is ignored"): + model.forward(batch, lead_times=lead_times, generator=generator) + # The warning is only emitted on the first offending forward call. + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + model.forward(batch, lead_times=lead_times, generator=generator) + assert not any("`generator` is ignored" in str(w.message) for w in caught) + assert torch.equal(state, generator.get_state()) + + +def test_rollout_passes_generator_through(): + model = _make_small_v1p5(stochastic=True) + model.eval() + generator = torch.Generator().manual_seed(42) + + def record_rollout(): + batch = _make_batch(surf_vars=_INPUT_SURF_VARS) + recorded = [] + original = model.backbone._sample_noise + + def recording_sample_noise(shape, device, dtype, generator=None): + noise = original(shape, device, dtype, generator) + recorded.append(noise.clone()) + return noise + + model.backbone._sample_noise = recording_sample_noise + try: + with torch.inference_mode(): + for _ in rollout(model, batch, steps=2, generator=generator): + pass + finally: + model.backbone._sample_noise = original + return recorded + + first = record_rollout() + generator.manual_seed(42) + model.reset_noise() + second = record_rollout() + assert len(first) == 2 + assert _seqs_equal(first, second) + + +def test_dtype_change_invalidates_noise_cache(): + model = _make_small_v1p5(stochastic=True) + model.eval() + model.set_noise_accumulation(n=2) + _record_noise(model, n_forwards=1) + model.double() + with pytest.warns(UserWarning, match="clearing noise cache"): + _record_noise(model, n_forwards=1) + model.set_noise_accumulation(n=0) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="Requires CUDA.") +def test_indexless_cuda_generator_accepted_on_cuda_model(): + model = _make_small_v1p5(stochastic=True).cuda() + model.eval() + # `torch.Generator(device="cuda")` reports device `cuda` without an index, which must be + # treated as compatible with the model's `cuda:0`, like PyTorch does. + generator = (torch.Generator(device="cuda").manual_seed(42),) + recorded = _record_noise(model, n_forwards=1, generator=generator) + assert recorded[0].device.type == "cuda" + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="Requires CUDA.") +def test_device_mismatch_raises_without_consuming_rng(): + model = _make_small_v1p5(stochastic=True).cuda() + model.eval() + generators = (torch.Generator().manual_seed(1), torch.Generator().manual_seed(2)) + states = [g.get_state().clone() for g in generators] + batch = _make_batch(surf_vars=_INPUT_SURF_VARS, batch_size=2) + with torch.inference_mode(), pytest.raises(ValueError, match="on device"): + model.forward(batch, lead_times=torch.full((2,), 6.0), generator=generators) + for state, g in zip(states, generators): + assert torch.equal(state, g.get_state()) From 9c996352e72453aad5b41c369f2ffcd2fe72a11b Mon Sep 17 00:00:00 2001 From: s-sasaki-earthsea-wizard Date: Sun, 13 Sep 2026 11:57:20 +0900 Subject: [PATCH 2/6] Simplify generator validation and deduplicate documentation Address review feedback on #202: - Drop the warning when a generator is passed to a non-stochastic model; the argument is now silently ignored and documented as such. - Document the generator semantics on Aurora.forward only and cross-reference it everywhere else. - Remove _validate_generator: keep only the tuple length check, inlined in the backbone forward, and let PyTorch raise on device mismatches. - Revert the noise cache invalidation to the shape-only check. - Export NoiseGenerator from swin3d.py via __all__. - Condense the docs example and match the surrounding style. --- aurora/model/aurora.py | 32 ++++-------------- aurora/model/swin3d.py | 77 +++++++----------------------------------- aurora/rollout.py | 8 ++--- docs/models.md | 33 ++++++++---------- 4 files changed, 35 insertions(+), 115 deletions(-) diff --git a/aurora/model/aurora.py b/aurora/model/aurora.py index edd185e9..4db2a09d 100644 --- a/aurora/model/aurora.py +++ b/aurora/model/aurora.py @@ -323,9 +323,6 @@ def __init__( if isinstance(m, (WindowAttention, PerceiverAttention)): m.use_fp16_safe_attention = True - # Warn only once when `generator` is passed to `forward` of a non-stochastic model. - self._generator_ignored_warned = False - def reset_noise(self) -> None: """Flush the backbone noise cache. @@ -357,33 +354,16 @@ def forward( `(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): Source of randomness for the noise injection in stochastic mode. A - single generator drives one stream for the whole batch. A tuple must contain one - entry per batch element (ensemble member), in the current order of the batch - dimension; every element then draws from its own stream, so the noise sequence of - a given member does not depend on the batch composition. Tuple entries may be - `None` to fall back to the global RNG for that member, and passing the same - generator object in several slots makes those members share one stream. Because - the two modes draw with different shapes, a single generator and a tuple are not - interchangeable. Generators must live on the same device as the model, and they - advance on every forward pass. To reproduce a run, re-seed the generators (e.g. - with `manual_seed`) *and* call :meth:`reset_noise`, so that noise cached by noise - accumulation in a previous run cannot contaminate the reproduced sequence. When - the model is not stochastic, this argument is ignored with a warning. Defaults to - `None`, which draws from the global RNG (the previous behaviour). + 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. """ - if generator is not None and not self.backbone.stochastic: - if not self._generator_ignored_warned: - warnings.warn( - "`generator` is ignored because stochastic noise is disabled.", - stacklevel=2, - ) - self._generator_ignored_warned = True - generator = None - batch = self.batch_transform_hook(batch) # Get the first parameter. We'll derive the data type and device from this parameter. diff --git a/aurora/model/swin3d.py b/aurora/model/swin3d.py index 7bace233..58d7a540 100644 --- a/aurora/model/swin3d.py +++ b/aurora/model/swin3d.py @@ -26,12 +26,9 @@ maybe_adjust_windows, ) -__all__ = ["Swin3DTransformerBackbone"] +__all__ = ["Swin3DTransformerBackbone", "NoiseGenerator"] NoiseGenerator = torch.Generator | tuple[torch.Generator | None, ...] | None -"""Source of randomness for noise injection in stochastic mode: a single generator driving one -stream for the whole batch, a tuple with one generator per batch element, or `None` for the -global RNG.""" class MLP(nn.Module): @@ -927,12 +924,6 @@ def reset_noise(self) -> None: Call this to clear all cached noise tensors, e.g. at the beginning of a new forecast issue time. After reset, the next forward call starts building the cache afresh. - - To reproduce a run that controls the noise with `generator` (see :meth:`forward`), re-seed - the generators *and* call this method: cached noise left over from a previous run would - otherwise contaminate the reproduced sequence. The same applies after changing the order or - composition of the ensemble members that a tuple of generators corresponds to, which cannot - be detected from the noise tensors themselves. """ if self.stochastic: @@ -957,35 +948,6 @@ 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 _validate_generator( - self, generator: NoiseGenerator, batch_size: int, device: torch.device - ) -> None: - """Validate `generator` against the batch before any randomness is consumed. - - Checking the tuple length and every device up front ensures that no generator has already - advanced when an error is raised for a later batch element. - """ - if not isinstance(generator, tuple): - return - if len(generator) != batch_size: - raise ValueError( - f"Expected {batch_size} generators (one per batch element), got {len(generator)}." - ) - for i, g in enumerate(generator): - if g is None: - continue - # Like PyTorch, treat an index-less device (e.g. `cuda`) as compatible with an - # indexed one (e.g. `cuda:0`); only compare indices when both are explicit. - if g.device.type != device.type or ( - g.device.index is not None - and device.index is not None - and g.device.index != device.index - ): - raise ValueError( - f"Generator for batch element {i} is on device `{g.device}`, but noise is " - f"generated on device `{device}`." - ) - def _sample_noise( self, shape: tuple[int, ...], @@ -993,15 +955,7 @@ def _sample_noise( dtype: torch.dtype, generator: NoiseGenerator = None, ) -> torch.Tensor: - """Draw one noise sample of shape `(B, L, D)`. - - A single generator produces the whole sample in one draw, so its stream depends on the - batch size. A tuple of generators instead produces one `(L, D)` draw per batch element, so - the sequence of draws for a given element does not depend on the batch composition. `None`, - or a `None` entry in a tuple, falls back to the global RNG. Because the two modes consume - a generator with differently shaped draws, a single generator and a tuple are not - interchangeable. - """ + """Draw noise of shape `shape`, one draw per batch element if `generator` is a tuple.""" if isinstance(generator, tuple): return torch.stack( [torch.randn(shape[1:], device=device, dtype=dtype, generator=g) for g in generator] @@ -1039,10 +993,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): Source of - randomness for noise injection in stochastic mode. See :meth:`_sample_noise` and - :meth:`Aurora.forward` for the semantics. Only used when the model is stochastic. - Defaults to `None`, which draws from the global RNG. + 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)`. @@ -1064,21 +1016,18 @@ def forward( if self.stochastic: noise_shape = x.shape[:-1] + (self.embed_dim,) - self._validate_generator(generator, x.shape[0], x.device) + if isinstance(generator, tuple) and len(generator) != x.shape[0]: + raise ValueError( + f"Expected one generator per batch element, but got `{len(generator)}` " + f"generators for a batch of size `{x.shape[0]}`." + ) noise = self._sample_noise(noise_shape, x.device, x.dtype, generator) if self._accumulate_noise: - # A shape (e.g. different batch size), device, or dtype change invalidates the - # cache. - cached = self._noise_cache[0] if self._noise_cache else None - if cached is not None and ( - cached.shape != noise.shape - or cached.device != noise.device - or cached.dtype != noise.dtype - ): + # Shape change (e.g. different batch size) invalidates the cache. + if self._noise_cache and self._noise_cache[0].shape != noise.shape: warnings.warn( - f"Cached noise of shape {cached.shape} ({cached.dtype} on " - f"{cached.device}) is incompatible with new noise of shape {noise.shape} " - f"({noise.dtype} on {noise.device}); clearing noise cache.", + f"Noise shape changed from {self._noise_cache[0].shape} to " + f"{noise.shape}; clearing noise cache.", stacklevel=2, ) self._noise_cache.clear() diff --git a/aurora/rollout.py b/aurora/rollout.py index ef13cc25..bcd133d4 100644 --- a/aurora/rollout.py +++ b/aurora/rollout.py @@ -93,12 +93,8 @@ def rollout( 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): Source of randomness for the noise injection of stochastic models, passed - to every `forward` call of the roll-out. See :meth:`aurora.Aurora.forward` for the - semantics. The generators advance on every (sub-)step, and a tuple corresponds to the - order of the batch dimension throughout the roll-out. To reproduce a roll-out, re-seed - the generators and call `model.reset_noise()` before starting. Default: `None`, which - draws from the global RNG. + 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. diff --git a/docs/models.md b/docs/models.md index fddb2faf..e0613717 100644 --- a/docs/models.md +++ b/docs/models.md @@ -453,34 +453,29 @@ noise at each sub-step instead, though this is not recommended. ### Reproducible Noise -By default, the injected noise is drawn from the global RNG, so the noise of an individual -ensemble member cannot easily be reproduced. To control the noise, pass a `torch.Generator` -to `Aurora.forward` or `rollout`. The generator must live on the same device as the model: +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 -import torch - -from aurora import rollout - device = next(model.parameters()).device generator = torch.Generator(device=device).manual_seed(42) -preds = [pred for pred in rollout(model, batch, steps=10, generator=generator)] + +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 (ensemble member) to control the noise of every member -separately. Every member then draws from its own stream, so a member's noise sequence does not -depend on the batch composition. Tuple entries may be `None` to fall back to the global RNG for -that member. Note that a single generator and a tuple of generators draw with different shapes -and are therefore not interchangeable. +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)) -preds = [pred for pred in rollout(model, batch, steps=10, generator=generators)] + +with torch.inference_mode(): + preds = [pred.to("cpu") for pred in rollout(model, batch, steps=4, generator=generators)] ``` -Generators advance on every forward pass. -To reproduce a run, re-seed the generators with `manual_seed` *and* call `model.reset_noise()`: -the latter clears noise cached by noise accumulation, which would otherwise contaminate the -reproduced sequence. +To reproduce a run, re-seed the generators and call `model.reset_noise()`, which clears noise +cached by noise accumulation. From 1ecf11689dc14367e0f76620152dbe17ad342367 Mon Sep 17 00:00:00 2001 From: s-sasaki-earthsea-wizard Date: Sun, 13 Sep 2026 11:57:27 +0900 Subject: [PATCH 3/6] Condense generator tests into four parametrised tests Adopt the structure suggested in review: a stochastic-model fixture, a reusable noise recorder that takes a callable, single-generator behaviour parametrised over the noise accumulation size, and one test each for tuples, length mismatches, and roll-outs. Drop the tests for code removed in the previous commit. --- tests/v1p5/test_forward_generator.py | 286 +++++++-------------------- 1 file changed, 70 insertions(+), 216 deletions(-) diff --git a/tests/v1p5/test_forward_generator.py b/tests/v1p5/test_forward_generator.py index aea562dc..2334a233 100644 --- a/tests/v1p5/test_forward_generator.py +++ b/tests/v1p5/test_forward_generator.py @@ -1,11 +1,8 @@ """Copyright (c) Microsoft Corporation. Licensed under the MIT license. -Tests for the `generator` argument of `Aurora.forward`, which controls the noise injection of -stochastic Aurora 1.5 models (issue #191). +Tests for the `generator` argument of `Aurora.forward`. """ -import warnings - import pytest import torch @@ -15,240 +12,97 @@ _INPUT_SURF_VARS = tuple(v for v in _SURF_VARS if v not in _OUTPUT_ONLY_SURF) -def _record_noise(model, n_forwards=3, generator=None, batch_size=1): - """Run `n_forwards` identical forward passes and return the noise samples drawn.""" - batch = _make_batch(surf_vars=_INPUT_SURF_VARS, batch_size=batch_size) - lead_times = torch.full((batch_size,), 6.0) - recorded = [] - original = model.backbone._sample_noise - - def recording_sample_noise(shape, device, dtype, generator=None): - noise = original(shape, device, dtype, generator) - recorded.append(noise.clone()) - return noise - - model.backbone._sample_noise = recording_sample_noise - try: - with torch.inference_mode(): - for _ in range(n_forwards): - model.forward(batch, lead_times=lead_times, generator=generator) - finally: - model.backbone._sample_noise = original - return recorded - - -def _seqs_equal(a, b): - return len(a) == len(b) and all(torch.equal(x, y) for x, y in zip(a, b)) - - -def _member_seq(recorded, i): - return [noise[i] for noise in recorded] - - -def test_single_generator_reseed_reproduces_noise(): - model = _make_small_v1p5(stochastic=True) - model.eval() - generator = torch.Generator().manual_seed(42) - first = _record_noise(model, generator=generator) - generator.manual_seed(42) - second = _record_noise(model, generator=generator) - assert _seqs_equal(first, second) - - -def test_single_generator_advances_without_reseed(): +@pytest.fixture +def model(): model = _make_small_v1p5(stochastic=True) model.eval() - generator = torch.Generator().manual_seed(42) - first = _record_noise(model, generator=generator) - second = _record_noise(model, generator=generator) - assert not _seqs_equal(first, second) - - -def test_same_seed_fresh_instance_reproduces_noise(): - # The core requirement of issue #191: a given ensemble member reproduces the same noise - # sequence across model runs, here even across fresh model instances. - model_a = _make_small_v1p5(stochastic=True) - model_b = _make_small_v1p5(stochastic=True) - model_a.eval() - model_b.eval() - generator_a = torch.Generator().manual_seed(42) - generator_b = torch.Generator().manual_seed(42) - assert _seqs_equal( - _record_noise(model_a, generator=generator_a), - _record_noise(model_b, generator=generator_b), - ) - - -def test_single_generator_independent_of_global_rng(): - model = _make_small_v1p5(stochastic=True) - model.eval() - generator = torch.Generator().manual_seed(42) - torch.manual_seed(0) - first = _record_noise(model, generator=generator) - generator.manual_seed(42) - torch.manual_seed(999) - second = _record_noise(model, generator=generator) - assert _seqs_equal(first, second) + return model -def test_generator_none_preserves_global_rng_semantics(): - model = _make_small_v1p5(stochastic=True) - model.eval() - torch.manual_seed(0) - first = _record_noise(model) - torch.manual_seed(0) - second = _record_noise(model) - torch.manual_seed(999) - third = _record_noise(model) - assert _seqs_equal(first, second) - assert not _seqs_equal(first, third) +def _record_noise(model, run): + """Call `run` and return the noise samples drawn.""" + recorded = [] + sample_noise = model.backbone._sample_noise + def record(*args): + recorded.append(sample_noise(*args)) + return recorded[-1] -def test_tuple_member_noise_independent_of_batch_composition(): - model = _make_small_v1p5(stochastic=True) - model.eval() - generators = tuple(torch.Generator().manual_seed(seed) for seed in (1, 2, 3)) - joint = _record_noise(model, generator=generators, batch_size=3) - alone = _record_noise(model, generator=(torch.Generator().manual_seed(3),), batch_size=1) - assert _seqs_equal(_member_seq(joint, 2), _member_seq(alone, 0)) + model.backbone._sample_noise = record + with torch.inference_mode(): + run() + model.backbone._sample_noise = sample_noise + return recorded -def test_tuple_and_single_generator_are_not_interchangeable(): - model = _make_small_v1p5(stochastic=True) - model.eval() - single = _record_noise(model, generator=torch.Generator().manual_seed(42), batch_size=2) - generators = (torch.Generator().manual_seed(42), torch.Generator().manual_seed(42)) - per_member = _record_noise(model, generator=generators, batch_size=2) - # Identically seeded per-member generators give every member the same noise, whereas a single - # generator drives one stream across the whole batch. - assert _seqs_equal(_member_seq(per_member, 0), _member_seq(per_member, 1)) - assert not _seqs_equal(_member_seq(single, 0), _member_seq(single, 1)) - assert not _seqs_equal(_member_seq(single, 1), _member_seq(per_member, 1)) +def _forward_noise(model, generator, batch_size=1, n_forwards=3): + """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 test_tuple_none_member_independent_of_other_generators(): - model = _make_small_v1p5(stochastic=True) - model.eval() - torch.manual_seed(123) - first = _record_noise(model, generator=(torch.Generator().manual_seed(42), None), batch_size=2) - torch.manual_seed(123) - second = _record_noise(model, generator=(torch.Generator().manual_seed(7), None), batch_size=2) - # The seeded member changes with its seed, but the `None` member draws from the global RNG - # and must not be affected by the other member's generator. - assert not _seqs_equal(_member_seq(first, 0), _member_seq(second, 0)) - assert _seqs_equal(_member_seq(first, 1), _member_seq(second, 1)) +def _equal(a, b): + return len(a) == len(b) and all(torch.equal(x, y) for x, y in zip(a, b)) -def test_tuple_length_mismatch_raises_without_consuming_rng(): - model = _make_small_v1p5(stochastic=True) - model.eval() - generators = tuple(torch.Generator().manual_seed(seed) for seed in (1, 2, 3)) - states = [g.get_state().clone() for g in generators] - batch = _make_batch(surf_vars=_INPUT_SURF_VARS, batch_size=2) - with torch.inference_mode(), pytest.raises(ValueError, match="Expected 2 generators"): - model.forward(batch, lead_times=torch.full((2,), 6.0), generator=generators) - for state, g in zip(states, generators): - assert torch.equal(state, g.get_state()) +def _member(recorded, i): + return [noise[i] for noise in recorded] -def test_accumulation_reproduces_after_reseed_and_reset_noise(): - model = _make_small_v1p5(stochastic=True) - model.eval() - model.set_noise_accumulation(n=2) +@pytest.mark.parametrize("n", [0, 2]) +def test_single_generator(model, n): + model.set_noise_accumulation(n) generator = torch.Generator().manual_seed(42) - first = _record_noise(model, generator=generator) - # Reproducing a run requires *both* re-seeding the generator and flushing the noise cache. + 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() - second = _record_noise(model, generator=generator) - # Re-seeding alone is not enough: leftover cached noise changes how many samples are drawn. - generator.manual_seed(42) - third = _record_noise(model, generator=generator) - model.set_noise_accumulation(n=0) - assert _seqs_equal(first, second) - assert not _seqs_equal(first, third) + torch.manual_seed(1) + assert _equal(first, _forward_noise(model, generator)) + # Without re-seeding, the generator keeps advancing. + assert not _equal(first, _forward_noise(model, generator)) + + +def test_tuple_of_generators(model): + def forward_noise(seeds, batch_size): + 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(_member(first, 0), _member(second, 0)) + assert _equal(_member(first, 2), _member(second, 2)) + assert _equal(_member(first, 2), _member(alone, 0)) + # ... and a `None` entry uses the global RNG. + assert _equal(_member(first, 1), _member(second, 1)) + + +def test_tuple_length_mismatch(model): + 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(),)) -def test_non_stochastic_model_warns_once_and_ignores_generator(): - model = _make_small_v1p5() - model.eval() +def test_rollout(model): generator = torch.Generator().manual_seed(42) - state = generator.get_state().clone() batch = _make_batch(surf_vars=_INPUT_SURF_VARS) - lead_times = torch.tensor([6.0]) - with torch.inference_mode(): - with pytest.warns(UserWarning, match="`generator` is ignored"): - model.forward(batch, lead_times=lead_times, generator=generator) - # The warning is only emitted on the first offending forward call. - with warnings.catch_warnings(record=True) as caught: - warnings.simplefilter("always") - model.forward(batch, lead_times=lead_times, generator=generator) - assert not any("`generator` is ignored" in str(w.message) for w in caught) - assert torch.equal(state, generator.get_state()) + def run(): + list(rollout(model, batch, steps=2, generator=generator)) -def test_rollout_passes_generator_through(): - model = _make_small_v1p5(stochastic=True) - model.eval() - generator = torch.Generator().manual_seed(42) - - def record_rollout(): - batch = _make_batch(surf_vars=_INPUT_SURF_VARS) - recorded = [] - original = model.backbone._sample_noise - - def recording_sample_noise(shape, device, dtype, generator=None): - noise = original(shape, device, dtype, generator) - recorded.append(noise.clone()) - return noise - - model.backbone._sample_noise = recording_sample_noise - try: - with torch.inference_mode(): - for _ in rollout(model, batch, steps=2, generator=generator): - pass - finally: - model.backbone._sample_noise = original - return recorded - - first = record_rollout() + first = _record_noise(model, run) generator.manual_seed(42) - model.reset_noise() - second = record_rollout() assert len(first) == 2 - assert _seqs_equal(first, second) - - -def test_dtype_change_invalidates_noise_cache(): - model = _make_small_v1p5(stochastic=True) - model.eval() - model.set_noise_accumulation(n=2) - _record_noise(model, n_forwards=1) - model.double() - with pytest.warns(UserWarning, match="clearing noise cache"): - _record_noise(model, n_forwards=1) - model.set_noise_accumulation(n=0) - - -@pytest.mark.skipif(not torch.cuda.is_available(), reason="Requires CUDA.") -def test_indexless_cuda_generator_accepted_on_cuda_model(): - model = _make_small_v1p5(stochastic=True).cuda() - model.eval() - # `torch.Generator(device="cuda")` reports device `cuda` without an index, which must be - # treated as compatible with the model's `cuda:0`, like PyTorch does. - generator = (torch.Generator(device="cuda").manual_seed(42),) - recorded = _record_noise(model, n_forwards=1, generator=generator) - assert recorded[0].device.type == "cuda" - - -@pytest.mark.skipif(not torch.cuda.is_available(), reason="Requires CUDA.") -def test_device_mismatch_raises_without_consuming_rng(): - model = _make_small_v1p5(stochastic=True).cuda() - model.eval() - generators = (torch.Generator().manual_seed(1), torch.Generator().manual_seed(2)) - states = [g.get_state().clone() for g in generators] - batch = _make_batch(surf_vars=_INPUT_SURF_VARS, batch_size=2) - with torch.inference_mode(), pytest.raises(ValueError, match="on device"): - model.forward(batch, lead_times=torch.full((2,), 6.0), generator=generators) - for state, g in zip(states, generators): - assert torch.equal(state, g.get_state()) + assert _equal(first, _record_noise(model, run)) From 365df602c0d349094f3ca863a8d7d92fde3fdf26 Mon Sep 17 00:00:00 2001 From: s-sasaki-earthsea-wizard Date: Tue, 15 Sep 2026 11:48:06 +0900 Subject: [PATCH 4/6] Move the generator tuple length check into _sample_noise Address review feedback on #202: the tuple length is what makes the stacked noise come out as the requested shape, so check it where it matters and drop the duplicate check in the backbone forward. --- aurora/model/swin3d.py | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/aurora/model/swin3d.py b/aurora/model/swin3d.py index 58d7a540..46f59805 100644 --- a/aurora/model/swin3d.py +++ b/aurora/model/swin3d.py @@ -957,6 +957,11 @@ def _sample_noise( ) -> 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] ) @@ -1016,11 +1021,6 @@ def forward( if self.stochastic: noise_shape = x.shape[:-1] + (self.embed_dim,) - if isinstance(generator, tuple) and len(generator) != x.shape[0]: - raise ValueError( - f"Expected one generator per batch element, but got `{len(generator)}` " - f"generators for a batch of size `{x.shape[0]}`." - ) 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. From b34397964286182444e662ffcf7b28702d58ca10 Mon Sep 17 00:00:00 2001 From: s-sasaki-earthsea-wizard Date: Tue, 15 Sep 2026 11:48:09 +0900 Subject: [PATCH 5/6] Type the generator tests and cover sub-stepped roll-outs Address review feedback on #202: - Annotate the helpers and tests, and reduce the fixture to a one-liner now that Module.eval returns the module. - Rename _equal to _equal_sequences_of_tensors to say what it compares. - Flush the noise cache before the third run in test_single_generator. With accumulation the first run draws an extra sample to fill the cache, so the final assertion held on length alone and never compared any values. - Parametrise test_rollout over fine_lead_times to cover both call sites that pass the generator, and add a third run to check that the generator keeps advancing. The number of draws depends on how the cache is filled when sub-stepping, so the length check goes. --- tests/v1p5/test_forward_generator.py | 61 ++++++++++++++++------------ 1 file changed, 34 insertions(+), 27 deletions(-) diff --git a/tests/v1p5/test_forward_generator.py b/tests/v1p5/test_forward_generator.py index 2334a233..647c06dd 100644 --- a/tests/v1p5/test_forward_generator.py +++ b/tests/v1p5/test_forward_generator.py @@ -1,30 +1,31 @@ """Copyright (c) Microsoft Corporation. Licensed under the MIT license. -Tests for the `generator` argument of `Aurora.forward`. +Tests for the `generator` argument of `Aurora.forward` and `rollout`. """ +from typing import Any, Callable, Sequence + import pytest import torch from ._helpers import _OUTPUT_ONLY_SURF, _SURF_VARS, _make_batch, _make_small_v1p5 -from aurora import rollout +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(): - model = _make_small_v1p5(stochastic=True) - model.eval() - return model +def model() -> AuroraV1p5: + return _make_small_v1p5(stochastic=True).eval() -def _record_noise(model, run): +def _record_noise(model: AuroraV1p5, run: Callable[[], Any]) -> list[torch.Tensor]: """Call `run` and return the noise samples drawn.""" - recorded = [] + recorded: list[torch.Tensor] = [] sample_noise = model.backbone._sample_noise - def record(*args): + def record(*args: Any) -> torch.Tensor: recorded.append(sample_noise(*args)) return recorded[-1] @@ -35,7 +36,9 @@ def record(*args): return recorded -def _forward_noise(model, generator, batch_size=1, n_forwards=3): +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) @@ -48,16 +51,16 @@ def _forward_noise(model, generator, batch_size=1, n_forwards=3): ) -def _equal(a, b): +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, i): +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, n): +def test_single_generator(model: AuroraV1p5, n: int) -> None: model.set_noise_accumulation(n) generator = torch.Generator().manual_seed(42) torch.manual_seed(0) @@ -67,13 +70,14 @@ def test_single_generator(model, n): generator.manual_seed(42) model.reset_noise() torch.manual_seed(1) - assert _equal(first, _forward_noise(model, generator)) + assert _equal_sequences_of_tensors(first, _forward_noise(model, generator)) # Without re-seeding, the generator keeps advancing. - assert not _equal(first, _forward_noise(model, generator)) + model.reset_noise() + assert not _equal_sequences_of_tensors(first, _forward_noise(model, generator)) -def test_tuple_of_generators(model): - def forward_noise(seeds, batch_size): +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) @@ -82,27 +86,30 @@ def forward_noise(seeds, batch_size): 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(_member(first, 0), _member(second, 0)) - assert _equal(_member(first, 2), _member(second, 2)) - assert _equal(_member(first, 2), _member(alone, 0)) + 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(_member(first, 1), _member(second, 1)) + assert _equal_sequences_of_tensors(_member(first, 1), _member(second, 1)) -def test_tuple_length_mismatch(model): +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(),)) -def test_rollout(model): +@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(): - list(rollout(model, batch, steps=2, generator=generator)) + 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 len(first) == 2 - assert _equal(first, _record_noise(model, run)) + 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)) From 8f0a98e3deaf3b06b37c5ed70e2bedc7b7d5966e Mon Sep 17 00:00:00 2001 From: s-sasaki-earthsea-wizard Date: Fri, 25 Sep 2026 18:49:51 +0900 Subject: [PATCH 6/6] Fix mypy errors reported by the pre-commit hook Declare NoiseGenerator with an explicit TypeAlias. The hook's environment does not install torch, so mypy could not recognise the assignment as an implicit type alias. In the generator tests, patch _sample_noise with patch.object instead of assigning to the method, which also restores it if the run raises. --- aurora/model/swin3d.py | 4 ++-- tests/v1p5/test_forward_generator.py | 5 ++--- 2 files changed, 4 insertions(+), 5 deletions(-) diff --git a/aurora/model/swin3d.py b/aurora/model/swin3d.py index 46f59805..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 @@ -28,7 +28,7 @@ __all__ = ["Swin3DTransformerBackbone", "NoiseGenerator"] -NoiseGenerator = torch.Generator | tuple[torch.Generator | None, ...] | None +NoiseGenerator: TypeAlias = torch.Generator | tuple[torch.Generator | None, ...] | None class MLP(nn.Module): diff --git a/tests/v1p5/test_forward_generator.py b/tests/v1p5/test_forward_generator.py index 647c06dd..288ed094 100644 --- a/tests/v1p5/test_forward_generator.py +++ b/tests/v1p5/test_forward_generator.py @@ -4,6 +4,7 @@ """ from typing import Any, Callable, Sequence +from unittest.mock import patch import pytest import torch @@ -29,10 +30,8 @@ def record(*args: Any) -> torch.Tensor: recorded.append(sample_noise(*args)) return recorded[-1] - model.backbone._sample_noise = record - with torch.inference_mode(): + with torch.inference_mode(), patch.object(model.backbone, "_sample_noise", record): run() - model.backbone._sample_noise = sample_noise return recorded