Save and load PINNs and other scimba_jax objects

Scimba PINNs can be saved to and loaded from files, allowing to:

  • save time and energy by avoiding to train again a PINN

  • resume the training of a PINN,

  • use a PINN trained on another machine or another device.

We also provide facilities for saving and loading other scimba_jax instances - to be more precise, all instances of classes inheriting from ScimbaPytree.

Save and load utilities are based on the orbax library; keep in mind that it saves and loads only the leaves of pytrees.

Let us first demonstrate how to save and load a trained PINN.

Save a trained PINN

We first define and train a PINN to approximate the solution of an Helmholtz problem.

[1]:
import jax
import jax.numpy as jnp
import matplotlib.pyplot as plt

import scimba_jax
from scimba_jax.domains.meshless_domains.domains_2d import Square2D
from scimba_jax.nonlinear_approximation.approximation_spaces.approximation_spaces import (
    ApproximationSpace,
)
from scimba_jax.nonlinear_approximation.integration.monte_carlo import (
    DomainSampler,
    TensorizedSampler,
)

from scimba_jax.nonlinear_approximation.networks.mlp import MLP
from scimba_jax.nonlinear_approximation.numerical_solvers.projectors import Projector
from scimba_jax.physical_models.elliptic_pde.helmholtz import HelmholtzND
from scimba_jax.plots.plots_nd import plot_abstract_approx_space

scimba_jax.set_verbosity(True)

/////////////// Scimba jax 0.0.1 ////////////////
Scimba_jax uses device: [CpuDevice(id=0)]
Scimba_jax uses dtype: <class 'jax.numpy.float64'>


[2]:
sigma = 0.03
center = [0.0, 0.5]

def f_rhs(xy: jnp.ndarray) -> jnp.ndarray:
    x, y = xy[0:1], xy[1:2]
    return (
        -100 * jnp.exp(-((x - center[0]) ** 2 + (y - center[1]) ** 2) / (2.0 * sigma**2.0)) / ((2.0 * jnp.pi) * sigma),
        jnp.zeros_like(x),
    )

def post_processing(approx: jnp.ndarray, xy: jnp.ndarray) -> jnp.ndarray:
    x, y = xy[0:1], xy[1:2]
    satisfy_dirichlet = approx * (x + 1.0) * (1.0 - x) * (y + 1.0) * (1.0 - y)
    satisfy_pure_real = jnp.stack([satisfy_dirichlet[...,0], jnp.zeros_like(satisfy_dirichlet[...,0])], axis=-1)
    return satisfy_pure_real

domain_x = [(-1.0, 1.0), (-1.0, 1.0)]

dx = Square2D(domain_x, is_main_domain=True)
sampler = TensorizedSampler([DomainSampler(dx)], bc=False)

model = HelmholtzND(
    dx, model_type="x", k_sq_re = (4*jnp.pi)**2, f_rhs=f_rhs, bc="strong_dirichlet", mode="single"
)

key = jax.random.PRNGKey(0)

nn = MLP(in_size=2, out_size=1, hidden_sizes=[20, 20], key=key,
         embedding="fourier",
         n_fourier_features=20,
         fourier_features_std=1.,
         embedding_axes=[0, 1],
        )
space = ApproximationSpace(
    {"x": 2}, [(nn, "vec", 2)], model_type="x", post_processing=post_processing
)

pinn = Projector(model, space, sampler)

key, pinn = pinn.project(key, space, 100, 8000)
Training: 100%|||||||||||||||||| 100/100[00:53<00:00] , loss: 1.5e+02 -> 6.4e-03
[3]:
plot_abstract_approx_space(
    pinn.space,
    dx,
    loss=pinn.losses,
    components=[0,],
    draw_contours=True,
    n_drawn_contours=20,
)
plt.show()
../_images/tutorials_jax_save_load_3_0.png

To save our trained PINN, we use the save method of the class Projector.

[4]:
pinn.save("tutorial")
Saving Projector to /Users/imbach/.scimba/scimba_jax/tutorial
Projector saved successfully
/Users/imbach/Work/new_scimba/src/scimba_jax/utils/scimba_pytree.py:436: UserWarning: Directory /Users/imbach/.scimba/scimba_jax/tutorial already exists and will be erased.
  warnings.warn(

By default, PINNs (and other scimba_jax objects) are saved in ~/.scimba/scimba_jax/.

If the directory already exists, it is overwritten with a warning.

[5]:
from scimba_jax.utils.paths import get_dirpath_for_save
dirpath_for_save = get_dirpath_for_save("tutorial")
import os
assert os.path.abspath(dirpath_for_save)

We will see below how to change the location of the saved PINN.

Let us first show how to load this PINN; one need to create a PINN like the one that has been saved, then load it from the file that has just been created.

[6]:
nn2 = MLP(in_size=2, out_size=1, hidden_sizes=[20, 20], key=key,
         embedding="fourier",
         n_fourier_features=20,
         fourier_features_std=1.,
         embedding_axes=[0, 1],
        )
space2 = ApproximationSpace(
    {"x": 2}, [(nn2, "vec", 2)], model_type="x", post_processing=post_processing
)

pinn2 = Projector(model, space2, sampler)

n_epochs, pinn2 = pinn2.load("tutorial")
print("n_epochs: ", n_epochs)
Loading Projector metadata from /Users/imbach/.scimba/scimba_jax/tutorial
Projector metadata loaded successfully
Loading Projector from /Users/imbach/.scimba/scimba_jax/tutorial
Projector loaded successfully
n_epochs:  100

n_epochs is the number of optimization epochs used to train the loaded PINN.

[7]:
plot_abstract_approx_space(
    pinn2.space,
    dx,
    loss=pinn2.losses,
    components=[0,],
    draw_contours=True,
    n_drawn_contours=20,
)
plt.show()
../_images/tutorials_jax_save_load_11_0.png

For the sake of reproducibility of trainings, the random number generator key is saved and loaded with projectors:

[8]:
assert jnp.all(pinn2.key == key)

When the saved PINN is loaded to a PINN which is not similar:

[9]:
pinn3 = Projector(model, space2, sampler, optimizer="SS-BFGS")
n_epochs, pinn3 = pinn3.load("tutorial")
Loading Projector metadata from /Users/imbach/.scimba/scimba_jax/tutorial
Projector metadata loaded successfully
Loading Projector from /Users/imbach/.scimba/scimba_jax/tutorial
/Users/imbach/Work/new_scimba/src/scimba_jax/utils/scimba_pytree.py:524: UserWarning: Incompatible structures for loading  from /Users/imbach/.scimba/scimba_jax/tutorial.
  warnings.warn(

In this case, the returned n_epochs is 0 and the returned projector is a copy of pinn3.

The behavior is similar when the directory to load from does not exist:

[10]:
pinn3 = Projector(model, space2, sampler, optimizer="SS-BFGS")
n_epochs, pinn3 = pinn3.load("tutorialtuto")
Loading Projector metadata from /Users/imbach/.scimba/scimba_jax/tutorialtuto
Directory /Users/imbach/.scimba/scimba_jax/tutorialtuto does not exist. Doing nothing.
Loading Projector from /Users/imbach/.scimba/scimba_jax/tutorialtuto
/Users/imbach/Work/new_scimba/src/scimba_jax/utils/scimba_pytree.py:511: UserWarning: Directory /Users/imbach/.scimba/scimba_jax/tutorialtuto does not exist. Doing nothing.
  warnings.warn(

The PINN can be saved and loaded on different machines or devices.

One can also continue the training of a PINN:

[11]:
key, pinn2 = pinn2.project(pinn2.key, pinn2.space, 100, 8000)
Training: 100%|||||||||||||||||| 100/100[01:16<00:00] , loss: 5.5e-03 -> 4.7e-04

Notice the use of pinn2.key as random number generator key for pinn2.project allowing reproducibility of training results with or without load.

[12]:
plot_abstract_approx_space(
    pinn2.space,
    dx,
    loss=pinn2.losses,
    components=[0,],
    draw_contours=True,
    n_drawn_contours=20,
)
plt.show()
../_images/tutorials_jax_save_load_21_0.png

save and load location

When scimba is verbose, information is displayed at save and load actions.

By default, scimba_jax objects are saved/loaded in the file ~/.scimba/scimba_jax/YOUR_NAME.pt where YOUR_NAME is the name you want to give to the file passed as first argument of the save/load methods.

YOUR_NAME can be the name of a script, for instance "tutorial.py" or the pythonvariable __file__; in this case, the extension will be ignored.

[13]:
pinn2.save("tutorial")
Saving Projector to /Users/imbach/.scimba/scimba_jax/tutorial
Projector saved successfully
/Users/imbach/Work/new_scimba/src/scimba_jax/utils/scimba_pytree.py:436: UserWarning: Directory /Users/imbach/.scimba/scimba_jax/tutorial already exists and will be erased.
  warnings.warn(
[14]:
pinn2.save("test/test/tutorial.py")
Saving Projector to /Users/imbach/.scimba/scimba_jax/tutorial
Projector saved successfully

One can also specify a post-fix for the filename:

[15]:
pinn2.save("tutorial", "postfixed")
_, _ = pinn2.load("tutorial", "postfixed")
Saving Projector to /Users/imbach/.scimba/scimba_jax/tutorial_postfixed
Projector saved successfully
Loading Projector metadata from /Users/imbach/.scimba/scimba_jax/tutorial_postfixed
Projector metadata loaded successfully
/Users/imbach/Work/new_scimba/src/scimba_jax/utils/scimba_pytree.py:436: UserWarning: Directory /Users/imbach/.scimba/scimba_jax/tutorial_postfixed already exists and will be erased.
  warnings.warn(
Loading Projector from /Users/imbach/.scimba/scimba_jax/tutorial_postfixed
Projector loaded successfully

and change the path of the destination directory:

[16]:
pinn2.save("tutorial", "retrained", path="~/saved_PINNS")
_, _ = pinn2.load("tutorial", "retrained", path="~/saved_PINNS")
Saving Projector to /Users/imbach/saved_PINNS/.scimba/scimba_jax/tutorial_retrained
Projector saved successfully
Loading Projector metadata from /Users/imbach/saved_PINNS/.scimba/scimba_jax/tutorial_retrained
Projector metadata loaded successfully
Loading Projector from /Users/imbach/saved_PINNS/.scimba/scimba_jax/tutorial_retrained
/Users/imbach/Work/new_scimba/src/scimba_jax/utils/scimba_pytree.py:436: UserWarning: Directory /Users/imbach/saved_PINNS/.scimba/scimba_jax/tutorial_retrained already exists and will be erased.
  warnings.warn(
Projector loaded successfully

or the all destination directory:

[17]:
pinn2.save("tutorial", "retrained", path="~", folder_name="scimba_PINNs")
_, _ = pinn2.load("tutorial", "retrained", path="~", folder_name="scimba_PINNs")
Saving Projector to /Users/imbach/scimba_PINNs/tutorial_retrained
Projector saved successfully
Loading Projector metadata from /Users/imbach/scimba_PINNs/tutorial_retrained
Projector metadata loaded successfully
Loading Projector from /Users/imbach/scimba_PINNs/tutorial_retrained
Projector loaded successfully
/Users/imbach/Work/new_scimba/src/scimba_jax/utils/scimba_pytree.py:436: UserWarning: Directory /Users/imbach/scimba_PINNs/tutorial_retrained already exists and will be erased.
  warnings.warn(

save and load for other scimba_jax objects

All instances of classes inheriting from ScimbaPytree have save and load methods, with behaviour similar to the one described above regarding saving directories location.

It can be interesting to save only a neural network:

[18]:
nn2 = pinn2.space.models[0]
assert isinstance(nn2, MLP)
nn2.save("tutorial", "NN")
Saving MLP to /Users/imbach/.scimba/scimba_jax/tutorial_NN
MLP saved successfully
/Users/imbach/Work/new_scimba/src/scimba_jax/utils/scimba_pytree.py:436: UserWarning: Directory /Users/imbach/.scimba/scimba_jax/tutorial_NN already exists and will be erased.
  warnings.warn(

It can only be loaded in a neural network with exactly the same structure.

[19]:
nn3 = MLP(in_size=2, out_size=1, hidden_sizes=[30, 30], key=key,
         embedding="fourier",
         n_fourier_features=20,
         fourier_features_std=1.,
         embedding_axes=[0, 1],
        )

success, nn3 = nn3.load("tutorial", "NN")
print("success: ", success)
Loading MLP from /Users/imbach/.scimba/scimba_jax/tutorial_NN
success:  False
/Users/imbach/Work/new_scimba/src/scimba_jax/utils/scimba_pytree.py:524: UserWarning: Incompatible structures for loading  from /Users/imbach/.scimba/scimba_jax/tutorial_NN.
  warnings.warn(
[20]:
nn3 = MLP(in_size=2, out_size=1, hidden_sizes=[20, 20], key=key,
         embedding="fourier",
         n_fourier_features=20,
         fourier_features_std=1.,
         embedding_axes=[0, 1],
        )

success, nn3 = nn3.load("tutorial", "NN")
print("success: ", success)
Loading MLP from /Users/imbach/.scimba/scimba_jax/tutorial_NN
MLP loaded successfully
success:  True

We finally embed it in an approximation space and plot it:

[21]:
space3 = ApproximationSpace(
    {"x": 2}, [(nn3, "vec", 2)], model_type="x", post_processing=post_processing
)

plot_abstract_approx_space(
    space3,
    dx,
    components=[0,],
    draw_contours=True,
    n_drawn_contours=20,
)
plt.show()
../_images/tutorials_jax_save_load_37_0.png
[ ]: