# Prediction w multiple shoot

**URL:** <https://discourse.julialang.org/t/prediction-w-multiple-shoot/101483>\
**Category:** New to Julia\
**Created:** [July 11, 2023, 1:36pm UTC](https://discourse.julialang.org/t/prediction-w-multiple-shoot/101483 "2023-07-11T13:36:52Z")\
**Posts on this page:** 6\
**Page:** 1

<div class="post-metadata">

**Author:** ![Marco\_Nesta](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/marco_nesta/32/50318_2.png) [@Marco\_Nesta](https://discourse.julialang.org/u/Marco_Nesta)\
**Post date:** [July 11, 2023, 1:36pm UTC](https://discourse.julialang.org/t/prediction-w-multiple-shoot/101483/1 "2023-07-11T13:36:52Z")

</div>

i am using multiple shoots to learn the lotka volterra equations. how can i predict results using the updated parameters?

emphasized textusing ComponentArrays, Lux, DiffEqFlux, Optimization, OptimizationPolyalgorithms, DifferentialEquations, Plots  
using DiffEqFlux: group\_ranges

using Random  
rng = Random.default\_rng()

# Define initial conditions and time steps

u0 = Float32[1.0 ; 1.0]  
tspan = (0.0f0, 10.0f0)  
datasize = 70

tsteps = range(tspan[1], tspan[2], length = datasize)  
function lotka\_volterra(du,u,p,t)  
x, y = u  
p = Float32[1.5;1.0;3.0;1.0]  
α, β, δ, γ = p  
du[1] = dx = α_x - β_x_y  
du[2] = dy = -δ_y + γ_x_y  
end

prob = ODEProblem(lotka\_volterra,u0,tspan)

# Verify ODE solution

ode\_data =Array(solve(prob, Tsit5(), saveat = tsteps))  
anim = Plots.Animation()

# Define the Neural Network

nn = Lux.Chain(x → x.^3,  
Lux.Dense(2, 84, swish),  
Lux.Dense(84, 44, swish),  
Lux.Dense(44, 22, swish),  
Lux.Dense(22, 12, swish),  
Lux.Dense(12,2))  
p\_init, st = Lux.setup(rng, nn)

neuralode = NeuralODE(nn, tspan, Tsit5(), saveat = tsteps)  
prob\_node = ODEProblem((u,p,t)-\>nn(u,p,st)[1], u0, tspan, ComponentArray(p\_init))

function plot\_multiple\_shoot(plt, preds, group\_size)  
step = group\_size-1  
ranges = group\_ranges(datasize, group\_size)

```
for (i, rg) in enumerate(ranges)
	plot!(plt, tsteps[rg], preds[i][1,:], markershape=:circle, label="Group $(i)")
end

```

end

# Animate training, cannot make animation on CI server

# anim = Plots.Animation()

iter = 0  
callback = function (p, l, preds; doplot = true)  
display(l)  
global iter  
iter += 1  
if doplot && iter%1 == 0  
# plot the original data  
plt = scatter(tsteps, ode\_data[1,:], label = “Data”)

```
# plot the different predictions for individual shoot
plot_multiple_shoot(plt, preds, group_size)

frame(anim,plt)
display(plot(plt))

```

end  
return false  
end

# Define parameters for Multiple Shooting

group\_size = 8  
continuity\_term = 200

function loss\_function(data, pred)  
return sum(abs2, data - pred)  
end

function loss\_multiple\_shooting(p)  
return multiple\_shoot(p, ode\_data, tsteps, prob\_node, loss\_function, Tsit5(),  
group\_size; continuity\_term)  
end

adtype = Optimization.AutoZygote()  
optf = Optimization.OptimizationFunction((x,p) → loss\_multiple\_shooting(x), adtype)  
optprob = Optimization.OptimizationProblem(optf, ComponentArray(p\_init))  
res\_ms = Optimization.solve(optprob, PolyOpt(),  
callback = callback)  
gif(anim, “multiple\_shooting.gif”, fps=15)

---

<div class="post-metadata">

**Author:** ![zdenek\_hurak](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/zdenek_hurak/32/53118_2.png) [@zdenek\_hurak](https://discourse.julialang.org/u/zdenek_hurak)\
**Post date:** [July 11, 2023, 2:16pm UTC](https://discourse.julialang.org/t/prediction-w-multiple-shoot/101483/2 "2023-07-11T14:16:07Z")

</div>

A minor formal issue: if you enclose your source code with triple backticks, the code will not only be better displayed but more convenient to copy and paste.

For example

```plaintext
x = 1:10

```

---

<div class="post-metadata">

**Author:** ![John\_Gibson](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/john_gibson/32/5321_2.png) [@John\_Gibson](https://discourse.julialang.org/u/John_Gibson)\
**Post date:** [July 11, 2023, 2:45pm UTC](https://discourse.julialang.org/t/prediction-w-multiple-shoot/101483/5 "2023-07-11T14:45:19Z")

</div>

Here’s your code properly indented and formatted for Markdown with triple backticks. There could be some editing errors in there. I did not try to run the code.

```julia
using ComponentArrays, Lux, DiffEqFlux, Optimization, OptimizationPolyalgorithms, DifferentialEquations, Plots
using DiffEqFlux: group_ranges

using Random
rng = Random.default_rng()

# Define initial conditions and time steps

u0 = Float32[1.0 ; 1.0]
tspan = (0.0f0, 10.0f0)
datasize = 70

tsteps = range(tspan[1], tspan[2], length = datasize)

function lotka_volterra(du,u,p,t)
    x, y = u
    p = Float32[1.5;1.0;3.0;1.0]
    α, β, δ, γ = p
    du[1] = dx = α*x - β*x<em>y
    du[2] = dy = -δ</em>y + γ*x*y
end

prob = ODEProblem(lotka_volterra,u0,tspan)

# Verify ODE solution

ode_data =Array(solve(prob, Tsit5(), saveat = tsteps))
anim = Plots.Animation()

# Define the Neural Network

nn = Lux.Chain(x → x.^3,
Lux.Dense(2, 84, swish),
Lux.Dense(84, 44, swish),
Lux.Dense(44, 22, swish),
Lux.Dense(22, 12, swish),
Lux.Dense(12,2))
p_init, st = Lux.setup(rng, nn)

neuralode = NeuralODE(nn, tspan, Tsit5(), saveat = tsteps)
prob_node = ODEProblem((u,p,t)->nn(u,p,st)[1], u0, tspan, ComponentArray(p_init))

function plot_multiple_shoot(plt, preds, group_size)
    step = group_size-1
    ranges = group_ranges(datasize, group_size)

    for (i, rg) in enumerate(ranges)
        plot!(plt, tsteps[rg], preds[i][1,:], markershape=:circle, label="Group $(i)")
    end
end

# Animate training, cannot make animation on CI server

# anim = Plots.Animation()

iter = 0
callback = function (p, l, preds; doplot = true)
    display(l)
    global iter
    iter += 1
    if doplot && iter%1 == 0
        # plot the original data
        plt = scatter(tsteps, ode_data[1,:], label = “Data”)

        # plot the different predictions for individual shoot
        plot_multiple_shoot(plt, preds, group_size)

        frame(anim,plt)
        display(plot(plt))
    end
    return false
end

# Define parameters for Multiple Shooting

group_size = 8
continuity_term = 200

function loss_function(data, pred)
    return sum(abs2, data - pred)
end

function loss_multiple_shooting(p)
    return multiple_shoot(p, ode_data, tsteps, prob_node, loss_function, Tsit5(), group_size; continuity_term)
end

adtype = Optimization.AutoZygote()
optf = Optimization.OptimizationFunction((x,p) → loss_multiple_shooting(x), adtype)
optprob = Optimization.OptimizationProblem(optf, ComponentArray(p_init))
res_ms = Optimization.solve(optprob, PolyOpt(),
callback = callback)
gif(anim, “multiple_shooting.gif”, fps=15)
```

---

<div class="post-metadata">

**Author:** ![Marco\_Nesta](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/marco_nesta/32/50318_2.png) [@Marco\_Nesta](https://discourse.julialang.org/u/Marco_Nesta)\
**Post date:** [July 11, 2023, 3:07pm UTC](https://discourse.julialang.org/t/prediction-w-multiple-shoot/101483/6 "2023-07-11T15:07:39Z")

</div>

The code works, I wanted to know how to predict future steps after training the model.

---

<div class="post-metadata">

**Author:** ![John\_Gibson](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/john_gibson/32/5321_2.png) [@John\_Gibson](https://discourse.julialang.org/u/John_Gibson)\
**Post date:** [July 11, 2023, 3:26pm UTC](https://discourse.julialang.org/t/prediction-w-multiple-shoot/101483/7 "2023-07-11T15:26:39Z")

</div>

Yes, I just reformatted your code so people could read and understand it, and so you could see how that’s done.

---

<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:** [July 11, 2023, 7:56pm UTC](https://discourse.julialang.org/t/prediction-w-multiple-shoot/101483/8 "2023-07-11T19:56:00Z")

</div>

Just do the same as the tutorials. `res_ms.u` is the learned parameters, so `remake(prob, p = res_ms.u)`. It’s no different than the `remake` in the callback.
