""" Use a DeepONet to learn the antiderivative operator from dynamical systems u'=f on [0,1] subject to u(0) = 0. The right-hand sides f_i are given as callables, and one also have the knowledge of pairs (x_i, f_i(x_i), F_i(x_i)). The deepOnet is trained both with physic and data. Both encoders and decoders are MLPs. The source terms are GRFs. This is inspired from example 4.1.1 in paper: DeepONet: Learning nonlinear operators for identifying differential equations based on the universal approximation theorem of operators https://arxiv.org/pdf/1910.03193 The deepOnet is parameterized by the grid_size, and its propagator p is defined as p(f) = f(grid) where grid is a regular grid of grid_size x values. The deepOnet is trained with a training batch of B pdes, each corresponding to a right-hand side f_i. The training is done with Adam and SS-BFGS. Each optimization epoch is done on a sub-batch of batch_size pdes. To assess the quality of the approximation, the mean of the relative L2 error is computed on the train set and on a test set of B other triples (x_grid, f_i(x_grid), F_i(x_grid)) """ from __future__ import annotations import timeit import jax import jax.numpy as jnp from scimba_jax.domains.meshless_domains.domains_1d import Segment1D from scimba_jax.neural_operator.physic_no.abstract_deeponet import AbstractDeepONet from scimba_jax.nonlinear_approximation.approximation_spaces.physic_no_approximation_spaces import ( # noqa: E501 PhysicNOApproximationSpace, ) from scimba_jax.nonlinear_approximation.integration.monte_carlo import ( DomainSampler, TensorizedSampler, ) from scimba_jax.nonlinear_approximation.numerical_solvers.physic_no_projectors import ( PhysicNOProjector, ) from scimba_jax.physical_models.abstract_physical_model import AbstractPhysicalModel from scimba_jax.physical_models.abstract_residuals import ( NDARRAY_TYPE, NDARRAYS_FUNC_TYPE, ) from scimba_jax.physical_models.ode.anti_derivative import ( AntiDerivative1D, make_grf_and_antiderivative, ) from scimba_jax.utils.functional_fields import ( AbstractFunctionalField, make_functional_field_class, ) F_TYPE = NDARRAYS_FUNC_TYPE | AbstractFunctionalField | None ##### Main parameters of the script B = 100 # the numer of models in training/testing batches N_EPOCHS_ADAM = 2000 N_EPOCHS_SSBFGS = 2000 batch_size = 50 # the size of model batches at each epoch grid_size = 100 # the size of the grid for data n_colloc = 1000 # the nb of collocation point for interior physical residual n_bc_colloc = 1 # the nb of collocation point for boundary physical residual n_dl_colloc = 1000 # the nb of collocation points for data residuals TRAIN_MODE_ADAM = "load" # possible values: "new", "load", "resume" TRAIN_MODE_SSBFGS = "load" # possible values: "new", "load", "resume" class DeepONetAntiDerivative(AbstractDeepONet): def __init__( self, grid_size: int = 100, encoder_size: int = 5, encoder_hidden_sizes: list[int] = [40, 40, 40], decoder_hidden_sizes: list[int] = [40, 40], activation: str = "tanh", ): self.grid = jnp.linspace(0, 1, grid_size) super().__init__( in_size=1, model_size=1, encoder_size=encoder_size, encoder_hidden_sizes=encoder_hidden_sizes, decoder_hidden_sizes=decoder_hidden_sizes, model_type="x", activation=activation, type_model="scalar", ) def pre_encoder_size(self) -> int: return jnp.size(self.grid, 0) def pre_encoder(self, physical_model: AbstractPhysicalModel) -> NDARRAY_TYPE: assert physical_model.main_domain is not None label = physical_model.main_domain.get_label() residual = physical_model.physical_residuals[label] jitted_evaluator = jax.jit(jax.vmap(lambda x: residual.f_rhs(x))) return jax.lax.stop_gradient(jitted_evaluator(self.grid)) def generate_funcs_pde_batch(key, nb_pde, x_eval_for_error): keys = jax.random.split(key, nb_pde + 1) key, subkeys = keys[0], keys[1:] fs_Fs = [make_grf_and_antiderivative(subkey) for subkey in subkeys] x_data = jnp.stack( [jax.random.uniform(subkey, shape=(n_dl_colloc, 1)) for subkey in subkeys], axis=0, ) y_data = jnp.stack( [jax.vmap(fs_Fs[i][1])(x_data[i, ...]) for i in range(nb_pde)], axis=0, ) fClass = make_functional_field_class("f") pdes = [ AntiDerivative1D( main_domain=domain_x, f_rhs=fClass(fs_Fs[i][0]), data=(x_data[i], y_data[i]), batchable_args=True, ) for i in range(nb_pde) ] exact_values = jnp.stack( [jax.vmap(fs_Fs[i][1])(x_eval_for_error) for i in range(nb_pde)], axis=0, ) return key, fs_Fs, pdes, exact_values @jax.jit def compute_mean_relative_l2_error( projector: PhysicNOProjector, batched_pdes: AbstractPhysicalModel, exact_values: NDARRAY_TYPE, x_eval: NDARRAY_TYPE, ) -> NDARRAY_TYPE: # Batch the PDEs batched_evaluate = jax.vmap(projector.evaluate, in_axes=(0, None)) values = batched_evaluate(batched_pdes, x_eval) # Compute relative L2 errors for each PDE in the batch l2_errors = jnp.linalg.norm(values - exact_values, axis=(1, 2)) l2_exacts = jnp.linalg.norm(exact_values, axis=(1, 2)) # Avoid division by zero relative_errors = jnp.where(l2_exacts > 0, l2_errors / l2_exacts, l2_errors) mean_relative_error = jnp.mean(relative_errors) return mean_relative_error key = jax.random.PRNGKey(0) domain_x = Segment1D((0.0, 1.0), is_main_domain=True) domain_x.set_boundaries_dict( { "bc zero": ["bc left"], } ) x_sampler = TensorizedSampler([DomainSampler(domain_x)], bc=True, model_type="x") x_eval_for_errors = jnp.linspace(0, 1, grid_size)[..., None] print("generating training GRFs and pdes") start = timeit.default_timer() key, fs_Fs, pdes, exact_values = generate_funcs_pde_batch(key, B, x_eval_for_errors) stop = timeit.default_timer() print("generating training GRFs and pdes OK in: ", stop - start, "\n\n") print("generating testing GRFs and pdes") start = timeit.default_timer() key, test_fs_Fs, test_pdes, test_exact_values = generate_funcs_pde_batch( key, B, x_eval_for_errors ) stop = timeit.default_timer() print("generating testing GRFs and pdes OK in: ", stop - start, "\n\n") no = DeepONetAntiDerivative(grid_size=grid_size) space = PhysicNOApproximationSpace( dims={"x": 1}, list_models=[no], model_type="x", ) custom_weights = {"interior": [1.0], "bc zero": [1.0], "data": [1.0]} print("generating projector") start = timeit.default_timer() projector = PhysicNOProjector( pdes, space, x_sampler, weights=custom_weights, ) stop = timeit.default_timer() print("generating projector OK in: ", stop - start, "\n\n") print("build_grad_loss_function") start = timeit.default_timer() grad_loss_func = projector.build_grad_loss_function() stop = timeit.default_timer() print("build_grad_loss_function OK in: ", stop - start, "\n\n") print("sample_physical_models") start = timeit.default_timer() key, batched_pdes, _ = projector.sample_physical_models(key, batch_size) stop = timeit.default_timer() print("sample_physical_models OK in: ", stop - start, "\n\n") key, dict_of_samples = projector.domain_sampler.sample_with_batched_pdes( key, batched_pdes, 10, 10, n_dl=n_dl_colloc ) # print(dict_of_samples) start = timeit.default_timer() grad_loss = grad_loss_func(space, dict_of_samples, batched_pdes) grad_loss.block_until_ready() stop = timeit.default_timer() print("\n\nFirst evaluation grad:", stop - start, "\n\n") print(" grad_loss.shape: ", grad_loss.shape) print(" grad_loss: ", grad_loss) print("\n\n@@@@@@@@@@@@@@@@@@@@ train with Adam @@@@@@@@@@@@@@@@") name = __file__ postfix = "Adam" bool_load = False if TRAIN_MODE_ADAM in ["resume", "load"]: n_epochs, projector = projector.load(name, postfix) bool_load = n_epochs > 0 if bool_load: key = projector.key bool_train = (TRAIN_MODE_ADAM in ["new", "resume"]) or (not bool_load) if bool_train: # train start = timeit.default_timer() key, projector = projector.project( key, projector.space, N_EPOCHS_ADAM, batch_size, n_colloc, n_bc_colloc, n_dl_colloc=n_dl_colloc, ) projector.best_loss["total"].block_until_ready() stop = timeit.default_timer() print("\n\nTime for %d optimization step:" % N_EPOCHS_ADAM, stop - start) # save projector.save(name, postfix) print("Best Loss: \n\n", projector.best_loss["total"]) print( "\n\n@@@@@@@@@@@@@@@@@@@@ Computing L2 errors after Adam training @@@@@@@@@@@@@@@@" ) # Evaluate on a subset of training batch (first batch_size elements) train_error = compute_mean_relative_l2_error( projector, AntiDerivative1D.create_batch(pdes), exact_values, x_eval_for_errors, ) print(f"Adam - Training batch relative L2 error: {train_error:.6e}") # Evaluate on test batch test_error = compute_mean_relative_l2_error( projector, AntiDerivative1D.create_batch(test_pdes), test_exact_values, x_eval_for_errors, ) print(f"Adam - Test batch relative L2 error: {test_error:.6e}") print("\n\n@@@@@@@@@@@@@@@@@@@@ train with SSBFGS @@@@@@@@@@@@@@@@") no2 = DeepONetAntiDerivative() space2 = PhysicNOApproximationSpace( dims={"x": 1}, list_models=[no2], model_type="x", ) print("space.ndof: ", space.ndof) projector2 = PhysicNOProjector( pdes, space2, x_sampler, weights=custom_weights, optimizer="SS-BFGS", ) name = __file__ postfix = "SSBFGS" bool_load = False if TRAIN_MODE_SSBFGS in ["resume", "load"]: n_epochs, projector2 = projector2.load(name, postfix) bool_load = n_epochs > 0 if bool_load: key = projector.key bool_train = (TRAIN_MODE_SSBFGS in ["new", "resume"]) or (not bool_load) if bool_train: # train start = timeit.default_timer() key, projector2 = projector2.project( key, projector2.space, N_EPOCHS_SSBFGS, batch_size, n_colloc, n_bc_colloc, n_dl_colloc=n_dl_colloc, ) projector2.best_loss["total"].block_until_ready() stop = timeit.default_timer() print("\n\nTime for %d optimization step:" % N_EPOCHS_SSBFGS, stop - start) # save projector.save(name, postfix) print("Best Loss: \n\n", projector2.best_loss["total"]) print( "\n\n@@@@@@@@@@@@@@@@@@@@ Computing L2 errors after SS-BFGS training @@@@@@@@@@@@@@@@" ) train_error = compute_mean_relative_l2_error( projector2, AntiDerivative1D.create_batch(pdes), exact_values, x_eval_for_errors, ) print(f"SS-BFGS - Training batch relative L2 error: {train_error:.6e}") # Evaluate on test batch test_error = compute_mean_relative_l2_error( projector2, AntiDerivative1D.create_batch(test_pdes), test_exact_values, x_eval_for_errors, ) print(f"SS-BFGS - Test batch relative L2 error: {test_error:.6e}") model = test_pdes[0] model0 = pdes[0] nF = jax.jit(jax.vmap(lambda x: test_fs_Fs[0][1](x[0]))) nF0 = jax.jit(jax.vmap(lambda x: fs_Fs[0][1](x[0]))) projector.plot( models=[model0, model], exact_sols=[nF0, nF], errors=[nF0, nF], equal_aspect=False, title="Physics informed DeepONet trained with Adam", titles=("pde of the train set", "pde of the test set"), ) projector2.plot( models=[model0, model], exact_sols=[nF0, nF], equal_aspect=False, errors=[nF0, nF], title="Physics informed DeepONet trained with SS-BFGS", titles=("pde of the train set", "pde of the test set"), )