# Save and load NeuralPDE model for postprocessing

**URL:** <https://discourse.julialang.org/t/save-and-load-neuralpde-model-for-postprocessing/114668>\
**Category:** Machine Learning\
**Tags:** question, pde, neural-network, save\
**Created:** [May 24, 2024, 9:16am UTC](https://discourse.julialang.org/t/save-and-load-neuralpde-model-for-postprocessing/114668 "2024-05-24T09:16:56Z")\
**Posts on this page:** 4\
**Page:** 1

<div class="post-metadata">

**Author:** ![PetrosStefanou](https://avatars.discourse-cdn.com/v4/letter/p/a88e57/32.png) [@PetrosStefanou](https://discourse.julialang.org/u/PetrosStefanou)\
**Post date:** [May 24, 2024, 9:16am UTC](https://discourse.julialang.org/t/save-and-load-neuralpde-model-for-postprocessing/114668/1 "2024-05-24T09:16:57Z")

</div>

Hi,

I want to save a trained NeuralPDE model and then load it to some other script at a later time for analysis and post-processing.

In [this reply](https://discourse.julialang.org/t/neualpde-model-save-load/70551/4) it is suggested to save the trained Lux parameters stored in `res.u`. I am doing that through

```julia
using JLD2
# @save "trained_model.jld2" res.u

```

However, the actual representation of the solution is not stored there and I am not sure where should I feed these parameters to get the trained model. Should one recreate the entire `sys, strategy, discretization, prob` sequence and then use the stored pretrained parameters? This is rather slow as it unnecessarily compiles a bunch of things that are not going to be used. The solution I am thinking is to save `discretization.phi` or `res.cache.f.f.phi`, bu can these be saved in the same way as `res.u`?

Thanks in advance for the help.

---

<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:** [June 15, 2024, 11:12am UTC](https://discourse.julialang.org/t/save-and-load-neuralpde-model-for-postprocessing/114668/2 "2024-06-15T11:12:57Z")

</div>

> [@PetrosStefanou](#):
>
> However, the actual representation of the solution is not stored there and I am not sure where should I feed these parameters to get the trained model.

You just re-evaluate the neural architecture.

> [@PetrosStefanou](#):
>
> Should one recreate the entire `sys, strategy, discretization, prob` sequence and then use the stored pretrained parameters?

You don’t need any of those to evaluate the trained model.

Take for example the starting tutorial on PDEs:

> **[Introduction to NeuralPDE for PDEs · NeuralPDE.jl](https://docs.sciml.ai/NeuralPDE/stable/tutorials/pdesystem/)**
>
> Documentation for NeuralPDE.jl.

The neural architecture is

```julia
chain = Lux.Chain(Dense(dim, 16, Lux.σ), Dense(16, 16, Lux.σ), Dense(16, 1))

```

We can initialize it the standard Lux way:

```julia
rng = Random.default_rng()
p, st = Lux.setup(rng, U)
const _st = st

```

`chain` is a function that takes `(x,p,st)`. Your trained parameters `p` is `res.u` from before. So you just need to evaluate it with `chain(x, res.u, _st)` and you’re good!

In total this code looks like:

```julia
using JLD2, Lux, Random, ComponentArrays
trained_p = load("example.jld2")["res.u"]

rng = Random.default_rng()
p, st = Lux.setup(rng, U)
const _st = st

# Evaluate at new points x
chain(x, trained_p, _st)

```

ComponentArrays is just needed because the saved `trained_p` will be a ComponentArray by default, though you could convert to a NamedTuple and that would remove that dep requirement.

---

<div class="post-metadata">

**Author:** ![PetrosStefanou](https://avatars.discourse-cdn.com/v4/letter/p/a88e57/32.png) [@PetrosStefanou](https://discourse.julialang.org/u/PetrosStefanou)\
**Post date:** [June 17, 2024, 7:53am UTC](https://discourse.julialang.org/t/save-and-load-neuralpde-model-for-postprocessing/114668/3 "2024-06-17T07:53:07Z")

</div>

Yes you are right. I was somehow confused and thought that you need `discretization.phi` in order to evaluate the solution. Thanks for the clarification.

By the way, why do you do `const _st = st`? Does it have any convenience or performance benefit?

---

<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:** [June 17, 2024, 9:10am UTC](https://discourse.julialang.org/t/save-and-load-neuralpde-model-for-postprocessing/114668/4 "2024-06-17T09:10:43Z")

</div>

> [@PetrosStefanou](#):
>
> By the way, why do you do `const _st = st`? Does it have any convenience or performance benefit?

There’s a tiny performance benefit and a tiny improvement to Enzyme because it improves the compiler’s ability to optimize since it’s constant. It probably doesn’t matter much but I just instinctually make sure code is always well-typed.
