# Training NN with Loss from Differential Equation

**URL:** <https://discourse.julialang.org/t/training-nn-with-loss-from-differential-equation/99669>\
**Category:** Machine Learning\
**Tags:** question, diffeq\
**Created:** [May 31, 2023, 12:09pm UTC](https://discourse.julialang.org/t/training-nn-with-loss-from-differential-equation/99669 "2023-05-31T12:09:56Z")\
**Posts on this page:** 1\
**Page:** 1

<div class="post-metadata">

**Author:** ![Duarte\_Magano](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/duarte_magano/32/50390_2.png) [@Duarte\_Magano](https://discourse.julialang.org/u/Duarte_Magano)\
**Post date:** [May 31, 2023, 12:09pm UTC](https://discourse.julialang.org/t/training-nn-with-loss-from-differential-equation/99669/1 "2023-05-31T12:09:56Z")

</div>

Hello!

Say I have a differential equation of the form _du/dt = NN(u)_, where _NN(u)_ is the output of a neural network that takes as input the current value of u.  
The boundary condition is _u(0) = 0_.  
I want to train the neural network in such a way that, for a specified time _t=2_ and position _uf=1_, the solution _u_ to the differential equation obeys _u(t)=uf_.  
This is meant to be a toy version of the problem of teaching an agent to find paths…

After reading about Flux and Differential equations, I figured that I could approach the problem with the following code

```plaintext
using Flux: train!, params
using DifferentialEquations
using SciMLSensitivity

# define NN model
model = Chain(
    Dense(1, 4, relu),
    Dense(4, 1, x -> σ.(x))
) 
ps = params(model)

# define differential equation
function f!(du, u, p, t) 
    du[1] = model([u[1]])[1] # NN controls velocity
end

# calculate final position for our NN controller
function final_position()
    u0 = [0.] # starts at 0.
    tspan = (0.0, 2.0) # systems evolves for time = 2.
    prob = ODEProblem(f!, u0, tspan) # set problem
    sol = solve(prob, Tsit5(), save_everystep = false, save_start = false) # solve diff equation
    Xf = sol.u[1][1] #final position
    return Xf
end

# define loss function
uf = 1.
function loss()
    Xf = final_position() # final position with NN controller starting at X0
    return (Xf - uf)^2
end

# set optimizer
opt = ADAM(0.3)

# define (empty) data
x_train = Iterators.repeated((), 100)

# do one training round
train!(loss, ps, x_train, opt)

```

But I got the following warning

 ![image](https://global.discourse-cdn.com/julialang/original/3X/a/4/a471e1dc2d1c325ff02697218daaa4dd3ed3e5a0.png)  
and then the parameters are not updated at all.

I guess that this is happening because it is having trouble auto-differentiating a function of the output of a differential equation…  
But I could not find a way to work around it.

Any help here would be immensely appreciated!  
Thanks in advance!

I also believe that I may not be using the appropriate (or up to date) Julia framework for this type of problem.  
How would you approach the problem?

Btw, the package versions are:

````julia
  [f6369f11] ForwardDiff v0.10.35
  [91a5bcdd] Plots v1.38.14
  [1ed8b502] SciMLSensitivity v7.32.0
  [90137ffa] StaticArrays v1.5.25
  [e88e6eb3] Zygote v0.6.61```
````
