Solving ODEs with Physics-Informed Neural Networks: https://
using NeuralPDE
using Lux
using OptimizationOptimisers
using OrdinaryDiffEq
using LinearAlgebra
using Random
using Plots
rng = Random.default_rng()
Random.seed!(rng, 42)Random.TaskLocalRNG()Solve ODEs¶
The true function:
model(u, p, t) = cospi(2t)model (generic function with 1 method)Prepare data
tspan = (0.0, 1.0)
u0 = 0.0
prob = ODEProblem(model, u0, tspan)ODEProblem with uType Float64 and tType Float64. In-place: false
Non-trivial mass matrix: false
timespan: (0.0, 1.0)
u0: 0.0Construct a neural network to solve the problem.
chain = Lux.Chain(Lux.Dense(1, 5, σ), Lux.Dense(5, 1))
ps, st = Lux.setup(rng, chain) |> Lux.f64((layer_1 = (weight = [1.0205187797546387; 0.4480646252632141; … ; -0.17203108966350555; 1.4574639797210693;;], bias = [-0.37966299057006836, 0.4062596559524536, -0.1814206838607788, 0.34669220447540283, -0.9501688480377197]), layer_2 = (weight = [-0.4852765202522278 0.1757996529340744 … 0.26079171895980835 -0.3895817995071411], bias = [-0.20251299440860748])), (layer_1 = NamedTuple(), layer_2 = NamedTuple()))Solve the ODE with NeuralPDE.NNODE().
optimizer = OptimizationOptimisers.Adam(0.1)
alg = NeuralPDE.NNODE(chain, optimizer, init_params = ps)
@time sol = solve(prob, alg, maxiters = 2000, saveat = 0.01, verbose = true)Fetching long content....
retcode: Success
Interpolation: Trained neural network interpolation
t: 0.0:0.01:1.0
u: 101-element Vector{Float64}:
0.0
0.009151421211917572
0.01825691104070158
0.027299171299466487
0.036259632563473705
0.04511843088055512
0.053854394660181805
0.06244504363196118
0.07086660188482889
0.07909402709230419
⋮
-0.07731899338316414
-0.06912985447088438
-0.060584330881115894
-0.05169195568108881
-0.04246213655204361
-0.03290414929420665
-0.023027132883315104
-0.012840085860246172
-0.0023518638577649416Comparing to the regular solver
sol2 = solve(prob, Tsit5(), saveat=sol.t)
plot(sol2, label = "Tsit5")
plot!(sol.t, sol.u, label = "NNODE")
This notebook was generated using Literate.jl.