""" Neural Operator Framework — FD solver + RBF decoder ==================================================== Architecture : IntegralEncoder(data) -> alpha=u_h + GaussianDecoder(alpha, x) -> u(x) Encoder : solves -D(x)·u'' + a(x)·u = f(x) with centered finite differences -> alpha = u_h ∈ R^N Decoder : u(x) = Σᵢ αᵢ · G(x - xᵢ) with G(r) = exp(-‖r‖²/2σ²) """ from __future__ import annotations import inspect import timeit # from abc import abstractmethod from typing import Callable import jax import jax.numpy as jnp # from jax.flatten_util import ravel_pytree from scimba_jax.domains.meshless_domains.base import VolumetricDomain from scimba_jax.domains.meshless_domains.domains_1d import Segment1D from scimba_jax.nonlinear_approximation.approximation_spaces.physic_no_approximation_spaces import ( AbstractPhysicNO, PhysicNOApproximationSpace, ) from scimba_jax.nonlinear_approximation.integration.monte_carlo import ( DomainSampler, TensorizedSampler, ) from scimba_jax.nonlinear_approximation.model_class.funcparam_vectorial import ( ParamScalarFunction, ) from scimba_jax.nonlinear_approximation.numerical_solvers.physic_no_projectors import ( PhysicNOProjector, ) from scimba_jax.physical_models.abstract_physical_model import ( PHYSICAL_RESIDUALS_TYPE, AbstractPhysicalModel, ) from scimba_jax.physical_models.abstract_residuals import ( F_RHS_TYPE, NDARRAY_TYPE, NDARRAYS_FUNC_TYPE, PARAM_FUNC_TYPE, InteriorResidual, ) from scimba_jax.physical_models.boundary_residuals import DirichletResidual from scimba_jax.utils.functional_fields import ( AbstractFunctionalField, make_functional_field_class, ) from scimba_jax.utils.scimba_pytree import dynamic, trainable def solve_reaction_diffusion_fd( grid: NDARRAY_TYPE, tau: NDARRAY_TYPE, D: Callable, a: Callable, f: Callable ) -> NDARRAY_TYPE: print(">>> TRACING solve_reaction_diffusion_fd") # dof shape (N,1) → each f(x) returns (1,) → .ravel() forces (N,) f_h = jax.vmap(f)(grid).ravel() # (N,) a_h = jax.vmap(a)(grid).ravel() # (N,) d_h = jax.vmap(D)(grid).ravel() # (N,) h = jnp.squeeze(grid[1] - grid[0]) # grid step # Line i : -D_i/h² u_{i-1} + (2D_i/h² + a_i) u_i - D_i/h² u_{i+1} = f_i diag_main = 2.0 * d_h / h**2 + tau * a_h # (N,) diag_lower = -d_h[1:] / h**2 # (N-1,) — under the diagonal diag_upper = -d_h[:-1] / h**2 # (N-1,) — over the diagonal A = ( jnp.diag(diag_main, 0) + jnp.diag(diag_lower, -1) + jnp.diag(diag_upper, 1) ) # (N, N) # print("A: ", A) return jnp.linalg.solve(A, f_h) # (N,) — values at inner nodes def _build_d_default(size: int) -> NDARRAYS_FUNC_TYPE: def d_one(*args: NDARRAY_TYPE) -> NDARRAY_TYPE: return jnp.ones((size,), dtype=jnp.float64) return d_one def _build_a_default(size: int) -> NDARRAYS_FUNC_TYPE: def a_zeros(*args: NDARRAY_TYPE) -> NDARRAY_TYPE: return jnp.zeros((size,), dtype=jnp.float64) return a_zeros F_TYPE = NDARRAYS_FUNC_TYPE | AbstractFunctionalField | None class ReactionDiffusion1DResidual(InteriorResidual): D: AbstractFunctionalField = dynamic() a: AbstractFunctionalField = dynamic() def __init__( self, domain: VolumetricDomain, D: F_TYPE = None, a: F_TYPE = None, f_rhs: F_TYPE = None, model_type="x", ): super().__init__(domain=domain, size=1, model_type=model_type, f_rhs=f_rhs) self.D = self._construct_d_a(D, _build_d_default) self.a = self._construct_d_a(a, _build_a_default) def _construct_d_a( self, func: F_TYPE, default: Callable ) -> AbstractFunctionalField: if isinstance(func, AbstractFunctionalField): return func else: functional_field_class = make_functional_field_class("auto") if func is None: return functional_field_class(default(1)) else: return functional_field_class(func) def construct_residual( self, *rho: PARAM_FUNC_TYPE, precomputed: dict | None = None ) -> PARAM_FUNC_TYPE: rho_ = rho[0] assert isinstance(rho_, ParamScalarFunction) lap = rho_.laplacian("x") def nd(x): return self.D(x) def na(x): return self.a(x) return -lap * nd + rho_ * na class ReactionDiffusionDirichlet1D(AbstractPhysicalModel): def __init__( self, main_domain: VolumetricDomain, D: F_TYPE = None, a: F_TYPE = None, f_rhs: F_RHS_TYPE = None, f_bc_rhs: F_RHS_TYPE = None, model_type="x", ): super().__init__(main_domain=main_domain) assert isinstance(self.main_domain, VolumetricDomain) self.physical_residuals: PHYSICAL_RESIDUALS_TYPE = { self.main_domain.get_label(): ReactionDiffusion1DResidual( domain=main_domain, D=D, a=a, f_rhs=f_rhs, model_type=model_type ), } for boundary in self.boundaries: self.physical_residuals[boundary] = DirichletResidual( domain=self.boundaries[boundary], model_type=model_type, f_rhs=f_bc_rhs, ) class DFsolver(AbstractPhysicNO): """ Pytree : fd_solver, sigma ∈ children (dynamiques, différentiables) rien en aux_data """ # ⚠ Explicitly False, and it HAS to be: an approximation space declares # its models trainable, so anything inside one inherits that status. # A field that must escape it says so -- declaring only the exceptions. grid: NDARRAY_TYPE = trainable(False) # the discretisation: data, batchable # ⚠ Declared, not selected by a hand-written partition: these two are the # only reals the optimiser may move here. tau: NDARRAY_TYPE = trainable(True) sigma: NDARRAY_TYPE = trainable(True) def __init__(self, grid: NDARRAY_TYPE, tau: float = 1.0, sigma: float = 1.0): super().__init__(model_type="x", model_size=1, type_model="scalar") self.grid = grid self.tau = jnp.array(tau, dtype=jnp.float64) self.sigma = jnp.array(sigma, dtype=jnp.float64) def ndof(self) -> int: return 2 def alpha_shape(self) -> tuple[int, ...]: return self.grid.shape def beta_shape(self) -> tuple[int, ...]: return self.alpha_shape() def encoder(self, physical_model: AbstractPhysicalModel) -> NDARRAY_TYPE: assert isinstance(physical_model, ReactionDiffusionDirichlet1D) assert physical_model.main_domain is not None label = physical_model.main_domain.get_label() residual = physical_model.physical_residuals[label] assert isinstance(residual, ReactionDiffusion1DResidual) jitted_solve = jax.jit(solve_reaction_diffusion_fd) return jitted_solve( grid=self.grid, tau=jnp.array(self.tau), D=residual.D, a=residual.a, f=residual.f_rhs, ) def decoder(self, beta: NDARRAY_TYPE, *args: NDARRAY_TYPE) -> NDARRAY_TYPE: # x : (d_x,), dof : (N, d_x) → diffs : (N, d_x) x = args[0] diffs = x - self.grid G = jnp.exp(-jnp.sum(diffs**2, axis=-1) / (2.0 * self.sigma**2)) # (N,) # print("decoder output shape: ", jnp.sum(alpha * G, keepdims=True).shape) return jnp.sum(beta * G, keepdims=True) # (1,) def propagator(self, alpha: NDARRAY_TYPE) -> NDARRAY_TYPE: return alpha def test_no_projector(N_dof=64, d_x=1, B_total=100, N_query=100, batch_size=50, d_u=1): print("\n\n######### ", inspect.currentframe().f_code.co_name) print("\n\n") domain_x = Segment1D((0.0, 1.0), is_main_domain=True) fClass = make_functional_field_class("f") DClass = make_functional_field_class("D") aClass = make_functional_field_class("a") pdes = [ ReactionDiffusionDirichlet1D( main_domain=domain_x, D=DClass(lambda x, i=i: jnp.exp(-i * x) + 0.1), # D=DClass(lambda x, i=i: test(x, i)), a=aClass(lambda x, i=i: 1.0 + 0.5 * jnp.cos(i * x)), # a=aClass(lambda x, i=i: 3.0 * i * jnp.ones_like(x)), f_rhs=fClass(lambda x, i=i: jnp.sin((i + 1) * jnp.pi * x)), # f_rhs=fClass(lambda x, i=i: i * jnp.ones_like(x)), ) for i in range(B_total) ] key = jax.random.PRNGKey(0) dof = jnp.linspace(0, 1, N_dof + 2)[1:-1].reshape(N_dof, d_x) model = DFsolver(grid=dof, tau=0.001, sigma=0.1) space = PhysicNOApproximationSpace( dims={"x": 1}, list_models=[model], model_type="x", ) space2 = PhysicNOApproximationSpace( dims={"x": 1}, list_models=[model], model_type="x", ) space3 = PhysicNOApproximationSpace( dims={"x": 1}, list_models=[model], model_type="x", ) x_sampler = TensorizedSampler([DomainSampler(domain_x)], bc=True, model_type="x") projector = PhysicNOProjector(pdes, space, x_sampler) loss_func = projector.build_losses_function() loss_func = jax.jit(loss_func) key, batched_pdes, _ = projector.sample_physical_models(key, batch_size) key, x_samples = projector.sample_physical_domain( key, batched_pdes, N_query, N_query ) start = timeit.default_timer() losses = loss_func(space, x_samples, batched_pdes) losses["interior"].block_until_ready() stop = timeit.default_timer() print("\n\nFirst evaluation:", stop - start, "\n\n") print("Losses: ", losses) key, batched_pdes, _ = projector.sample_physical_models(key, batch_size) key, x_samples = projector.sample_physical_domain( key, batched_pdes, N_query, N_query ) start = timeit.default_timer() losses = loss_func(space2, x_samples, batched_pdes) losses["interior"].block_until_ready() stop = timeit.default_timer() print("\n\nSecond evaluation:", stop - start, "\n\n") print("Losses: ", losses) key, batched_pdes, _ = projector.sample_physical_models(key, batch_size) key, x_samples = projector.sample_physical_domain( key, batched_pdes, N_query, N_query ) start = timeit.default_timer() losses = loss_func(space3, x_samples, batched_pdes) losses["interior"].block_until_ready() stop = timeit.default_timer() print("\n\nThird evaluation:", stop - start, "\n\n") print("Losses: ", losses) key, batched_pdes, _ = projector.sample_physical_models(key, batch_size) key, x_samples = projector.sample_physical_domain( key, batched_pdes, N_query, N_query ) grad_loss_func = projector.build_grad_loss_function() start = timeit.default_timer() grad_loss = grad_loss_func(space, x_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) key, batched_pdes, _ = projector.sample_physical_models(key, batch_size) key, x_samples = projector.sample_physical_domain( key, batched_pdes, N_query, N_query ) start = timeit.default_timer() grad_loss = grad_loss_func(space2, x_samples, batched_pdes) grad_loss.block_until_ready() stop = timeit.default_timer() print("\n\nSecond evaluation grad:", stop - start, "\n\n") print(" grad_loss.shape: ", grad_loss.shape) print(" grad_loss: ", grad_loss) one_step_optim = projector.build_one_step_optim(batch_size, N_query, N_query) nspace = space start = timeit.default_timer() loss, nspace, key, opt = one_step_optim(nspace, key, projector.optimizer) loss["total"].block_until_ready() stop = timeit.default_timer() print("\n\nFirst evaluation one_step_optim jitted:", stop - start, "\n\n") print("new loss: ", loss) start = timeit.default_timer() loss, nspace, key, opt = one_step_optim(nspace, key, opt) loss["total"].block_until_ready() stop = timeit.default_timer() print("\n\nSecond evaluation one_step_optim jitted:", stop - start, "\n\n") print("new loss: ", loss) start = timeit.default_timer() loss, nspace, key, opt = one_step_optim(nspace, key, opt) loss["total"].block_until_ready() stop = timeit.default_timer() print("\n\nThird evaluation one_step_optim jitted:", stop - start, "\n\n") print("new loss: ", loss) start = timeit.default_timer() loss, nspace, key, opt = one_step_optim(nspace, key, opt) loss["total"].block_until_ready() stop = timeit.default_timer() print("\n\nFourth evaluation one_step_optim jitted:", stop - start, "\n\n") print("new loss: ", loss) N_EPOCHS_ADAM = 1000 print("\n\n@@@@@@@@@@@@@@@@@@@@ train with Adam @@@@@@@@@@@@@@@@") start = timeit.default_timer() key, projector = projector.project( key, space, N_EPOCHS_ADAM, batch_size, N_query, N_query ) projector.best_loss["total"].block_until_ready() stop = timeit.default_timer() print("\n\nTime for %d optimization step:" % N_EPOCHS_ADAM, stop - start) print("Best Loss: \n\n", projector.best_loss["total"]) projector.plot([pdes[1], pdes[2]], equal_aspect=False) print("\n\n@@@@@@@@@@@@@@@@@@@@ train with SS-BFGS @@@@@@@@@@@@@@@@") N_EPOCHS_SSBFGS = 100 model = DFsolver(grid=dof, tau=0.001, sigma=0.1) space = PhysicNOApproximationSpace( dims={"x": 1}, list_models=[model], model_type="x", ) # jax.config.update("jax_log_compiles", True) projector = PhysicNOProjector( pdes, space, x_sampler, optimizer="SS-BFGS", ) one_step_optim = projector.build_one_step_optim(batch_size, N_query, N_query) nspace = space start = timeit.default_timer() loss, nspace, key, opt = one_step_optim(nspace, key, projector.optimizer) loss["total"].block_until_ready() stop = timeit.default_timer() print("\n\nFirst evaluation one_step_optim jitted:", stop - start, "\n\n") print("new loss: ", loss) start = timeit.default_timer() loss, nspace, key, opt = one_step_optim(nspace, key, opt) loss["total"].block_until_ready() stop = timeit.default_timer() print("\n\nSecond evaluation one_step_optim jitted:", stop - start, "\n\n") print("new loss: ", loss) start = timeit.default_timer() loss, nspace, key, opt = one_step_optim(nspace, key, opt) loss["total"].block_until_ready() stop = timeit.default_timer() print("\n\nThird evaluation one_step_optim jitted:", stop - start, "\n\n") print("new loss: ", loss) start = timeit.default_timer() loss, nspace, key, opt = one_step_optim(nspace, key, opt) loss["total"].block_until_ready() stop = timeit.default_timer() print("\n\nFourth evaluation one_step_optim jitted:", stop - start, "\n\n") print("new loss: ", loss) start = timeit.default_timer() key, projector = projector.project( key, space, N_EPOCHS_SSBFGS, batch_size, N_query, N_query ) projector.best_loss["total"].block_until_ready() stop = timeit.default_timer() print("\n\nTime for %d optimization step:" % N_EPOCHS_SSBFGS, stop - start) print("Best Loss: \n\n", projector.best_loss["total"]) projector.plot([pdes[1], pdes[2]], equal_aspect=False) # jax.config.update("jax_log_compiles", False) # # # jax.config.update("jax_log_compiles", True) test_no_projector()