# Reduce the batches of a Parallel Ensemble Problem to the mean of the square modulus

**URL:** <https://discourse.julialang.org/t/reduce-the-batches-of-a-parallel-ensemble-problem-to-the-mean-of-the-square-modulus/88359>\
**Category:** General Usage\
**Tags:** differentialequation\
**Created:** [October 6, 2022, 4:04pm UTC](https://discourse.julialang.org/t/reduce-the-batches-of-a-parallel-ensemble-problem-to-the-mean-of-the-square-modulus/88359 "2022-10-06T16:04:13Z")\
**Posts on this page:** 9\
**Page:** 1

<div class="post-metadata">

**Author:** ![albertomercurio](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/albertomercurio/32/27051_2.png) [@albertomercurio](https://discourse.julialang.org/u/albertomercurio)\
**Post date:** [October 6, 2022, 4:04pm UTC](https://discourse.julialang.org/t/reduce-the-batches-of-a-parallel-ensemble-problem-to-the-mean-of-the-square-modulus/88359/1 "2022-10-06T16:04:14Z")

</div>

Hi,

I’m performing a generic EnsembleProblem. Let’s take a simple example

```julia
prob = ODEProblem((u,p,t)->1.01u,0.5,(0.0,1.0))

function prob_func(prob,i,repeat)
  remake(prob,u0=rand()*prob.u0)
end

ensemble_prob = EnsembleProblem(prob,prob_func=prob_func)
sim = solve(ensemble_prob,Tsit5(),EnsembleDistributed(),trajectories=100,batch_size = 20)

```

and I want that every batch\_size steps it performs the `timeseries_steps_mean` of the square modulus of the soultions.

I can do it when the simulation is finished, by doing `timestep_mean(abs2.(sim), 1:step)`, but how can I perform this after every batch? I think that the `reduce` function is the way, but I didn’t found a working method.

---

<div class="post-metadata">

**Author:** ![ChrisRackauckas](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/chrisrackauckas/32/77_2.png) [@ChrisRackauckas](https://discourse.julialang.org/u/ChrisRackauckas)\
**Post date:** [October 6, 2022, 5:04pm UTC](https://discourse.julialang.org/t/reduce-the-batches-of-a-parallel-ensemble-problem-to-the-mean-of-the-square-modulus/88359/2 "2022-10-06T17:04:26Z")

</div>

> [@albertomercurio](#):
>
> I think that the `reduce` function is the way, but I didn’t found a working method.

In the reduce function, just sum up the result of the batch divided by the batch length.

---

<div class="post-metadata">

**Author:** ![albertomercurio](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/albertomercurio/32/27051_2.png) [@albertomercurio](https://discourse.julialang.org/u/albertomercurio)\
**Post date:** [October 6, 2022, 7:22pm UTC](https://discourse.julialang.org/t/reduce-the-batches-of-a-parallel-ensemble-problem-to-the-mean-of-the-square-modulus/88359/3 "2022-10-06T19:22:01Z")

</div>

I tried something like

```julia
prob = ODEProblem((u,p,t)->1.01*u, [0.5,0.5], (0.0,1.0))

function prob_func(prob,i,repeat)
  remake(prob,u0=rand().*prob.u0)
end

function reduction(u,data,I)
    (u+sum(abs2.(data)),false)
end

ensemble_prob = EnsembleProblem(prob, prob_func=prob_func, reduction=reduction)
sim = solve(ensemble_prob,Tsit5(),EnsembleSerial(),trajectories=100,batch_size = 20)

```

But it doesn’t work. It gives me an error.

---

<div class="post-metadata">

**Author:** ![ChrisRackauckas](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/chrisrackauckas/32/77_2.png) [@ChrisRackauckas](https://discourse.julialang.org/u/ChrisRackauckas)\
**Post date:** [October 6, 2022, 8:44pm UTC](https://discourse.julialang.org/t/reduce-the-batches-of-a-parallel-ensemble-problem-to-the-mean-of-the-square-modulus/88359/4 "2022-10-06T20:44:40Z")

</div>

what’s the error?

---

<div class="post-metadata">

**Author:** ![albertomercurio](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/albertomercurio/32/27051_2.png) [@albertomercurio](https://discourse.julialang.org/u/albertomercurio)\
**Post date:** [October 6, 2022, 9:43pm UTC](https://discourse.julialang.org/t/reduce-the-batches-of-a-parallel-ensemble-problem-to-the-mean-of-the-square-modulus/88359/5 "2022-10-06T21:43:28Z")

</div>

```julia
MethodError: no method matching abs2(::ODESolution{Float64, 2, Vector{Vector{Float64}}, Nothing, 
Nothing, Vector{Float64}, Vector{Vector{Vector{Float64}}}, ODEProblem{Vector{Float64}, 
Tuple{Float64, Float64}, false, SciMLBase.NullParameters, ODEFunction{false, 
SciMLBase.AutoSpecialize, var"#435#436", UniformScaling{Bool}, Nothing, Nothing, Nothing, 
Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, 
typeof(SciMLBase.DEFAULT_OBSERVED), Nothing, Nothing}, Base.Pairs{Symbol, Union{}, Tuple{},
 NamedTuple{(), Tuple{}}}, SciMLBase.StandardODEProblem}, 
Tsit5{typeof(OrdinaryDiffEq.trivial_limiter!), typeof(OrdinaryDiffEq.trivial_limiter!), Static.False}, 
OrdinaryDiffEq.InterpolationData{ODEFunction{false, SciMLBase.AutoSpecialize, var"#435#436", 
UniformScaling{Bool}, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, 
Nothing, Nothing, Nothing, Nothing, Nothing, typeof(SciMLBase.DEFAULT_OBSERVED), Nothing, 
Nothing}, Vector{Vector{Float64}}, Vector{Float64}, Vector{Vector{Vector{Float64}}}, 
OrdinaryDiffEq.Tsit5ConstantCache{Float64, Float64}}, DiffEqBase.DEStats})

```

---

<div class="post-metadata">

**Author:** ![ChrisRackauckas](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/chrisrackauckas/32/77_2.png) [@ChrisRackauckas](https://discourse.julialang.org/u/ChrisRackauckas)\
**Post date:** [October 7, 2022, 4:21am UTC](https://discourse.julialang.org/t/reduce-the-batches-of-a-parallel-ensemble-problem-to-the-mean-of-the-square-modulus/88359/6 "2022-10-07T04:21:08Z")

</div>

> [@albertomercurio](#):
>
> `no method matching abs2(::ODESolution`

`data` is a `Vector{ODESolution}`. You would need to broadcast that on each `ODESolution`. `(u+sum(map(abs2,map(abs2,data))),false)` is a straightforward way, but there are ways to make that nicer.

---

<div class="post-metadata">

**Author:** ![albertomercurio](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/albertomercurio/32/27051_2.png) [@albertomercurio](https://discourse.julialang.org/u/albertomercurio)\
**Post date:** [October 8, 2022, 4:38pm UTC](https://discourse.julialang.org/t/reduce-the-batches-of-a-parallel-ensemble-problem-to-the-mean-of-the-square-modulus/88359/7 "2022-10-08T16:38:31Z")

</div>

It gives me the error

```julia
MethodError: no method matching abs2(::ODESolution{Float64, 2, Vector{Vector{Float64}}, Nothing, 
Nothing, Vector{Float64}, Vector{Vector{Vector{Float64}}}, ODEProblem{Vector{Float64}, 
Tuple{Float64, Float64}, false, SciMLBase.NullParameters, ODEFunction{false, 
SciMLBase.AutoSpecialize, var"#11#12", UniformScaling{Bool}, Nothing, Nothing, Nothing, Nothing, 
Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, 
typeof(SciMLBase.DEFAULT_OBSERVED), Nothing, Nothing}, Base.Pairs{Symbol, Union{}, Tuple{}, NamedTuple{(), Tuple{}}}, SciMLBase.StandardODEProblem}, 
Tsit5{typeof(OrdinaryDiffEq.trivial_limiter!), typeof(OrdinaryDiffEq.trivial_limiter!), Static.False}, 
OrdinaryDiffEq.InterpolationData{ODEFunction{false, SciMLBase.AutoSpecialize, var"#11#12", 
UniformScaling{Bool}, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, 
Nothing, Nothing, Nothing, Nothing, Nothing, typeof(SciMLBase.DEFAULT_OBSERVED), Nothing, 
Nothing}, Vector{Vector{Float64}}, Vector{Float64}, Vector{Vector{Vector{Float64}}}, 
OrdinaryDiffEq.Tsit5ConstantCache{Float64, Float64}}, DiffEqBase.DEStats})
Closest candidates are:

```

I solved instead using the help of `output_func`:

```julia
prob = ODEProblem((u,p,t)->1.01*u, [0.5,0.5], (0.0,1.0))

function prob_func(prob,i,repeat)
  remake(prob,u0=rand().*prob.u0)
end

function output_func(sol, i)
  (hcat(map(x->abs2.(x), sol.u)...), false)
end

function reduction(u,batch,I)
  tmp = sum(cat(batch..., dims = 3), dims = 3)
  length(u) == 0 && return tmp, false
  cat(u, tmp, dims = 3), false
end

ensemble_prob = EnsembleProblem(prob, prob_func=prob_func, output_func=output_func, reduction=reduction)
sim = solve(ensemble_prob,Tsit5(),EnsembleSerial(),trajectories=100,batch_size = 20);
solution = sum(sim.u, dims = 3) ./ 100

```

---

<div class="post-metadata">

**Author:** ![ChrisRackauckas](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/chrisrackauckas/32/77_2.png) [@ChrisRackauckas](https://discourse.julialang.org/u/ChrisRackauckas)\
**Post date:** [October 8, 2022, 5:49pm UTC](https://discourse.julialang.org/t/reduce-the-batches-of-a-parallel-ensemble-problem-to-the-mean-of-the-square-modulus/88359/8 "2022-10-08T17:49:21Z")

</div>

Your last code there looks fine and runs fine?

---

<div class="post-metadata">

**Author:** ![albertomercurio](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/albertomercurio/32/27051_2.png) [@albertomercurio](https://discourse.julialang.org/u/albertomercurio)\
**Post date:** [October 8, 2022, 11:21pm UTC](https://discourse.julialang.org/t/reduce-the-batches-of-a-parallel-ensemble-problem-to-the-mean-of-the-square-modulus/88359/9 "2022-10-08T23:21:11Z")

</div>

Yes it works. Anyway, thank you for your help!
