Repository navigation
Internal ensembling in Aurora models - #199
Agnieszka Słowik (Slowika) wants to merge 22 commits into
Conversation
…for backward compatibility.
118564d to
f0fa368
Compare
|
Thanks Agnieszka Słowik (@Slowika) for opening a PR! I replied on the issue and mentioned the idea of using the batch size to produce multiple ensemble members simultaneously. Do you think that approach would suffice, or does this capability need to be added to the model explicitly? |
I've addressed your feedback. Should be ready for review now! Jonathan Weyn (@jweyn) |
Wessel (wesselb)
left a comment
There was a problem hiding this comment.
Thanks for the changes, Agnieszka Słowik (@Slowika)! This is looking much simpler now. Really great! :) I've left some Claude-assisted comments. After those, I think this is ready to be merged!
Co-authored-by: Wessel <wessel.p.bruinsma@gmail.com>
Co-authored-by: Wessel <wessel.p.bruinsma@gmail.com>
|
Hi Wessel (Wessel (@wesselb))! Thank you for the review. I addressed all comments and made some of the suggested changes. For the remaining Claude suggestions, they contradict your suggestions in the issue, especially regarding tiling. Could you please have another look? |
|
Thanks, Agnieszka Słowik (@Slowika)! I've replied to all outstanding comments. |
|
Hi Wessel (Wessel (@wesselb))! Thank you for the review :). I have addressed all comments and made the suggested changes. |
There was a problem hiding this comment.
Thanks, Agnieszka Słowik (@Slowika)! This is looking great now. A few more comments and suggestions. I'd be happy to merge this in after those. :)
snip: deleted Claude comment about the PR description
Co-authored-by: Wessel <wessel.p.bruinsma@gmail.com>
Co-authored-by: Wessel <wessel.p.bruinsma@gmail.com>
Co-authored-by: Wessel <wessel.p.bruinsma@gmail.com>
Co-authored-by: Wessel <wessel.p.bruinsma@gmail.com>
Co-authored-by: Wessel <wessel.p.bruinsma@gmail.com>
Co-authored-by: Wessel <wessel.p.bruinsma@gmail.com>
Co-authored-by: Wessel <wessel.p.bruinsma@gmail.com>
Co-authored-by: Wessel <wessel.p.bruinsma@gmail.com>
|
Thanks for the review, Wessel (@wesselb)! Your comments have been addressed and I agree with all of your suggestions. I just need one final re-approval before this PR can be merged. |
Addressed Issue #192.
Developed with the aid of AI, in line with the Code of Conduct.
Jonathan Weyn (@jweyn) Wessel (@wesselb)
Problem
Running an
N-member ensemble currently requires a loop that callsAurora.forward()/rollout()once per member, and then manually combining the results. This under-utilises the GPU (Nseparate launches) when the GPU is capable of storing all of the ensemble state.Change
Add a
num_ensemble_membersconstructor argument toAurora(default1, fully backwards compatible). When set toN > 1,forward()/rollout()run allNmembers as a single, fused batched computation internally, rather thanNseparate calls. This is useful when combined withstochastic=True: since the backbone's existing per-batch-element noise injection means every instance receives independent noise.This is purely an additional option: looping over
forward()/rollout()to implement ensembling remains fully supported and unaffected.Design notes
Batch's public shape contract is untouched: no new dimension, no new methods. The batch-dimension tiling used to fuse the computation is a private implementation detail (_tile_batch/_split_batchinaurora/batch.py), never exposed onBatchitself.forward()'s return type is nowBatch | list[Batch]: a plainBatchwhennum_ensemble_members == 1(no change from today), or alist[Batch]ofNstandard-shaped batches:pred[m]is memberm's ordinary, individually inspectableBatch.rollout()follows the same contract per yielded step, keeping the tiled representation internal across autoregressive steps for efficiency, and temporarily forcingmodel.num_ensemble_members = 1during its loop (restored viatry/finally, even on early generator closure) so nestedforward()calls don't re-tile.num_ensemble_members > 1is requested on a non-stochasticmodel, since all members would then be identical.Tests
Added
tests/v1p5/test_ensemble.pycovering:_tile_batch/_split_batchround-trippingforward()'s single-Batchvs.list[Batch]return contractstochastic=Truevs. identity understochastic=Falserollout()'s per-step output shape plusnum_ensemble_membersrestoration (including on early.close()).