@@ -10,15 +10,30 @@ function SciMLBase.__solve(
1010 };
1111 trajectories, batch_size = trajectories,
1212 unstable_check = (dt, u, p, t) -> false , adaptive = true ,
13+ seed = nothing ,
14+ rng = nothing ,
15+ rng_func = SciMLBase. default_rng_func,
1316 kwargs...
1417 )
1518 if trajectories == 1
1619 return SciMLBase. __solve (
1720 ensembleprob, alg, EnsembleSerial (); trajectories = 1 ,
18- kwargs...
21+ seed, rng, rng_func, kwargs...
1922 )
2023 end
2124
25+ # Pre-generate per-trajectory seeds for reproducibility (matching SciMLBase v3 protocol)
26+ sim_seeds = (rng != = nothing || seed != = nothing ) ?
27+ SciMLBase. generate_sim_seeds (rng, seed, trajectories) : nothing
28+
29+ # Bundle ensemble RNG state for passing to SciMLBase.solve_batch (CPU offload path)
30+ ensemble_rng_state = (;
31+ sim_seeds,
32+ _solve_rng_mode = Val (:none ),
33+ rng_func,
34+ master_rng = rng,
35+ )
36+
2237 cpu_trajectories = (
2338 (
2439 ensemblealg isa EnsembleGPUArray ||
@@ -53,8 +68,8 @@ function SciMLBase.__solve(
5368
5469 function f ()
5570 return SciMLBase. solve_batch (
56- ensembleprob, _alg, EnsembleThreads (), cpu_II, nothing ;
57- kwargs...
71+ ensembleprob, _alg, EnsembleThreads (), cpu_II, nothing ,
72+ ensemble_rng_state; kwargs...
5873 )
5974 end
6075
@@ -69,6 +84,7 @@ function SciMLBase.__solve(
6984 time = @elapsed sol = batch_solve (
7085 ensembleprob, alg, ensemblealg,
7186 1 : gpu_trajectories, adaptive;
87+ sim_seeds, rng_func, master_rng = rng,
7288 unstable_check = unstable_check, kwargs...
7389 )
7490 if cpu_trajectories != 0
@@ -83,6 +99,7 @@ function SciMLBase.__solve(
8399 similar (
84100 batch_solve (
85101 ensembleprob, alg, ensemblealg, 1 : batch_size, adaptive;
102+ sim_seeds, rng_func, master_rng = rng,
86103 unstable_check = unstable_check, kwargs...
87104 ),
88105 0
@@ -100,6 +117,7 @@ function SciMLBase.__solve(
100117 end
101118 batch_data = batch_solve (
102119 ensembleprob, alg, ensemblealg, I, adaptive;
120+ sim_seeds, rng_func, master_rng = rng,
103121 unstable_check = unstable_check, kwargs...
104122 )
105123 if ensembleprob. reduction != = SciMLBase. DEFAULT_REDUCTION
@@ -120,6 +138,7 @@ function SciMLBase.__solve(
120138 end
121139 x = batch_solve (
122140 ensembleprob, alg, ensemblealg, I, adaptive;
141+ sim_seeds, rng_func, master_rng = rng,
123142 unstable_check = unstable_check, kwargs...
124143 )
125144 yield ()
@@ -145,10 +164,20 @@ function SciMLBase.__solve(
145164 end
146165end
147166
167+ function _make_ensemble_context (i, sim_seeds, rng_func, master_rng)
168+ sim_seed = sim_seeds != = nothing ? sim_seeds[i] : nothing
169+ pre_ctx = SciMLBase. EnsembleContext (i, 1 , 0 , sim_seed, nothing , master_rng)
170+ sim_rng = rng_func (pre_ctx)
171+ return @set pre_ctx. rng = sim_rng
172+ end
173+
148174function batch_solve (
149175 ensembleprob, alg,
150176 ensemblealg:: Union{EnsembleArrayAlgorithm, EnsembleKernelAlgorithm} , I,
151177 adaptive;
178+ sim_seeds = nothing ,
179+ rng_func = SciMLBase. default_rng_func,
180+ master_rng = nothing ,
152181 kwargs...
153182 )
154183 @assert ! isempty (I)
@@ -157,17 +186,18 @@ function batch_solve(
157186 return if ensemblealg isa EnsembleGPUKernel
158187 if ensembleprob. safetycopy
159188 probs = map (I) do i
189+ ctx = _make_ensemble_context (i, sim_seeds, rng_func, master_rng)
160190 make_prob_compatible (
161191 ensembleprob. prob_func (
162192 deepcopy (ensembleprob. prob),
163- i,
164- 1
193+ ctx
165194 )
166195 )
167196 end
168197 else
169198 probs = map (I) do i
170- make_prob_compatible (ensembleprob. prob_func (ensembleprob. prob, i, 1 ))
199+ ctx = _make_ensemble_context (i, sim_seeds, rng_func, master_rng)
200+ make_prob_compatible (ensembleprob. prob_func (ensembleprob. prob, ctx))
171201 end
172202 end
173203 # Using inner saveat requires all of them to be of same size,
@@ -246,7 +276,7 @@ function batch_solve(
246276 ReturnCode. Terminated :
247277 ReturnCode. Success
248278 ),
249- i
279+ _make_ensemble_context (I[i], sim_seeds, rng_func, master_rng)
250280 )[1 ]
251281 end
252282 for i in eachindex (probs)
@@ -258,11 +288,13 @@ function batch_solve(
258288 else
259289 if ensembleprob. safetycopy
260290 probs = map (I) do i
261- ensembleprob. prob_func (deepcopy (ensembleprob. prob), i, 1 )
291+ ctx = _make_ensemble_context (i, sim_seeds, rng_func, master_rng)
292+ ensembleprob. prob_func (deepcopy (ensembleprob. prob), ctx)
262293 end
263294 else
264295 probs = map (I) do i
265- ensembleprob. prob_func (ensembleprob. prob, i, 1 )
296+ ctx = _make_ensemble_context (i, sim_seeds, rng_func, master_rng)
297+ ensembleprob. prob_func (ensembleprob. prob, ctx)
266298 end
267299 end
268300 u0 = reduce (hcat, Array (probs[i]. u0) for i in 1 : length (I))
@@ -316,7 +348,7 @@ function batch_solve(
316348 stats = sol. stats,
317349 retcode = sol. retcode
318350 ),
319- i
351+ _make_ensemble_context (I[i], sim_seeds, rng_func, master_rng)
320352 )[1 ]
321353 for i in 1 : length (probs)
322354 ]
@@ -339,7 +371,7 @@ function batch_solve(
339371 stats = sol. stats,
340372 retcode = sol. retcode
341373 ),
342- i
374+ _make_ensemble_context (I[i], sim_seeds, rng_func, master_rng)
343375 )[1 ]
344376 for i in 1 : length (probs)
345377 ]
@@ -539,13 +571,13 @@ function ChainRulesCore.rrule(
539571end
540572
541573function solve_batch (
542- prob, alg, ensemblealg:: EnsembleThreads , II, pmap_batch_size;
543- kwargs...
574+ prob, alg, ensemblealg:: EnsembleThreads , II, pmap_batch_size,
575+ ensemble_rng_state; kwargs...
544576 )
545577 if length (II) == 1 || Threads. nthreads () == 1
546578 return SciMLBase. solve_batch (
547- prob, alg, EnsembleSerial (), II, pmap_batch_size;
548- kwargs...
579+ prob, alg, EnsembleSerial (), II, pmap_batch_size,
580+ ensemble_rng_state; kwargs...
549581 )
550582 end
551583
@@ -565,8 +597,8 @@ function solve_batch(
565597 I_local = II[(batch_size * (i - 1 ) + 1 ): (batch_size * i)]
566598 end
567599 SciMLBase. solve_batch (
568- prob, alg, EnsembleSerial (), I_local, pmap_batch_size;
569- kwargs...
600+ prob, alg, EnsembleSerial (), I_local, pmap_batch_size,
601+ ensemble_rng_state; kwargs...
570602 )
571603 end
572604 return SciMLBase. tighten_container_eltype (batch_data)
0 commit comments