# Multiple shooting with multiple trajectories: how to properly implement this?

**URL:** <https://discourse.julialang.org/t/multiple-shooting-with-multiple-trajectories-how-to-properly-implement-this/118715>\
**Category:** Optimization (Mathematical)\
**Tags:** question, diffeq\
**Created:** [August 28, 2024, 2:25pm UTC](https://discourse.julialang.org/t/multiple-shooting-with-multiple-trajectories-how-to-properly-implement-this/118715 "2024-08-28T14:25:51Z")\
**Posts on this page:** 3\
**Page:** 1

<div class="post-metadata">

**Author:** ![PyJulia](https://avatars.discourse-cdn.com/v4/letter/p/2acd7d/32.png) [@PyJulia](https://discourse.julialang.org/u/PyJulia)\
**Post date:** [August 28, 2024, 2:25pm UTC](https://discourse.julialang.org/t/multiple-shooting-with-multiple-trajectories-how-to-properly-implement-this/118715/1 "2024-08-28T14:25:51Z")

</div>

[DiffEq](https://docs.sciml.ai/DiffEqFlux/stable/examples/multiple_shooting/) has a tutorial on multiple shooting for neural ode’s. It works for a single trajectory. I have a dataset with several trajectories, each trajectory being a solution to a differential equation with different initial condition. I have an implementation that runs but I doubt it is efficient or the best way of implementing it. Can you tell me whether there are packages with tutorials that demonstrate how to do this properly and/or show me how I can improve my code?

In general terms, I took the part of the tutorial

```julia
function loss_multiple_shooting(p)
    ps = ComponentArray(p, pax)
    return multiple_shoot(ps, 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, pd)
res_ms = Optimization.solve(optprob, PolyOpt(); callback = callback)

```

and replaced it with

```julia
function loss_multiple_shooting(p, p_axes, group_size, dataset, ode_problem)
    ps = ComponentArray(p, p_axes)
    return multiple_shoot(ps, dataset, time_steps, ode_problem, loss_function,
        time_stepper, group_size; continuity_term)
end

function solution_to_parameters(solution::Vector{Float32})::ModelParams
    reshaped = reshape(solution,(Nx,Nx+1))
    return (
        weight = reshaped[:,1:end-1], 
        bias = reshape(reshaped[:,end],(:,1))
    )
end

...

losses = []
for (group_size, epochs) in epochs_per_group_size
    for epoch = ProgressBar(1:epochs)
        loss = 0
        for batch in ProgressBar(train_loader)
            trajectory = batch[1]
            u0 = OffsetArrays.no_offset_view(trajectory[:, 1])

            ps = ComponentArray(parameters)
            p_data, p_axes = getdata(ps), getaxes(ps)
                        
            ode_problem = ODEProblem(
                (u, p, t) -> model(u, p, st)[1], 
                u0, 
                time_span, 
                ps
            )

            opt_func = Optimization.OptimizationFunction((x, p) -> loss_multiple_shooting(x, p_axes, group_size, trajectory, ode_problem), ad_type)
            opt_prob = Optimization.OptimizationProblem(opt_func, p_data)
            result_solve = Optimization.solve(opt_prob, OptimizationOptimisers.Adam(), maxiters = solver_iterations; callback = callback);

            global parameters = solution_to_parameters(result_solve.u)

            loss += result_solve.objective
        end
        push!(losses, loss)
        display((epoch, loss))
        flush(stdout)
    end
end

```

I do a lot of `ComponentArray` calling. This seems inefficient. Is there a better way of updating u0 for the ode\_problem and the parameters for the neural network model?

---

<div class="post-metadata">

**Author:** ![baggepinnen](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/baggepinnen/32/693_2.png) [@baggepinnen](https://discourse.julialang.org/u/baggepinnen)\
**Post date:** [August 28, 2024, 4:51pm UTC](https://discourse.julialang.org/t/multiple-shooting-with-multiple-trajectories-how-to-properly-implement-this/118715/2 "2024-08-28T16:51:59Z")

</div>

Hello and welcome to the community 👋

Have you read [the performance tips in the julia manual](https://docs.julialang.org/en/v1/manual/performance-tips/)? In particular, the first point

> [Performance critical code should be inside a function](https://docs.julialang.org/en/v1/manual/performance-tips/#Performance-critical-code-should-be-inside-a-function)

After that try to profile your code, e.g., using `@profview` and `@profview_allocs` in vscode, to learn what is taking time and get a feeling for where to possibly improve your code.

---

<div class="post-metadata">

**Author:** ![PyJulia](https://avatars.discourse-cdn.com/v4/letter/p/2acd7d/32.png) [@PyJulia](https://discourse.julialang.org/u/PyJulia)\
**Post date:** [September 3, 2024, 1:00pm UTC](https://discourse.julialang.org/t/multiple-shooting-with-multiple-trajectories-how-to-properly-implement-this/118715/3 "2024-09-03T13:00:25Z")

</div>

Hi Baggepinnen, thank you for the pointer. I wasn’t aware that I had to do so.

I placed it into a function like

```julia
function main_loop(model, parameters, epochs_per_group_size::Dict{Int, Int}, train_loader, solver_iterations::Int32)
    losses = []
    for (group_size, epochs) in epochs_per_group_size
        for epoch = ProgressBar(1:epochs)
            loss::Float32 = 0.0
            for batch in ProgressBar(train_loader)
                trajectory = batch[1]
                u0 = OffsetArrays.no_offset_view(trajectory[:, 1])

                ps = ComponentArray(parameters)
                p_data, p_axes = getdata(ps), getaxes(ps)

                ode_problem = ODEProblem(
                    (u, p, t) -> model(u, p, st)[1],
                    u0,
                    time_span, 
                    ps;
                    isoutofdomain = outofdomain
                )

                opt_func = Optimization.OptimizationFunction((x, p) -> loss_multiple_shooting(x, p_axes, group_size, trajectory, ode_problem), ad_type)
                opt_prob = Optimization.OptimizationProblem(opt_func, p_data)
                result_solve = Optimization.solve(opt_prob, OptimizationOptimisers.Adam(); callback = callback, maxiters = solver_iterations, progress = true);

                parameters = solution_to_parameters(result_solve.u)

                loss += result_solve.objective
            end
            push!(losses, loss)
            display((epoch, loss))
            flush(stdout)
        end
    end
    return losses, parameters
end

```

Running profview and profview\_allocs on them shows me some graphs, I do not completely understand. The graphs suggest most time/compute is spent on the solve function. In the meanwhile, I have found functions like `remake` and `ncycle` that seem to be alternative ways of doing part of the inner loop work. If someone could tell me whether using that actually matter, would help a lot. Otherwise, I’ll just have to see using some numerical tests.
