import time import jax import jax.numpy as jnp from scimba_jax.domains.meshless_domains.domains_1d import Segment1D from scimba_jax.linear_approximation.basis.analytic_bases import local_taylor_basis from scimba_jax.linear_approximation.basis.general_bases import ( PatchwiseParametricBasis, ) from scimba_jax.linear_approximation.galerkin.dg.elliptic_dg_scheme import ( EllipticDGscheme, ) from scimba_jax.linear_approximation.galerkin.dg.flux import SIPGFlux from scimba_jax.linear_approximation.meshes.mesh import Mesh from scimba_jax.linear_approximation.quad.gauss_quad import UnitSquareTensorized from scimba_jax.linear_approximation.variables.variables_dg import VariablesDG from scimba_jax.mapping.mapping import InvertibleFunction, Mapping 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.networks.mlp import MLP 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_physical_weak_model import ( AbstractPhysicalWeakModel, ) from scimba_jax.physical_models.abstract_residuals import ( NDARRAY_TYPE, ) from scimba_jax.physical_models.classical_weakform.laplacian_weak_form import ( LaplacianWeakForm, ) from scimba_jax.physical_models.elliptic_pde.laplacians import LaplacianDirichletDG from scimba_jax.utils.functional_fields import ( make_functional_field_class, ) from scimba_jax.utils.scimba_pytree import dynamic physical_dim = 1 out_dim = 1 seed = 1 def f_rhs(x: jnp.ndarray) -> jnp.ndarray: return (4 * jnp.pi**2.0) * jnp.sin(2.0 * jnp.pi * x[0:1]) def f_rhs2(x: jnp.ndarray) -> jnp.ndarray: return ((2.1**2) * jnp.pi**2.0) * jnp.sin(2.1 * jnp.pi * x[0:1]) def u_exact(x: jnp.ndarray) -> jnp.ndarray: return jnp.sin(2.0 * jnp.pi * x[..., 0:1]) def u_exact2(x: jnp.ndarray) -> jnp.ndarray: return jnp.sin(2.1 * jnp.pi * x[..., 0:1]) def dirichlet_bc(x: jnp.ndarray) -> jnp.ndarray: return jnp.zeros(out_dim) def make_mesh(n_cells, quad_order): mapping_id = InvertibleFunction(lambda x: x, lambda y: y) return Mesh( dim=physical_dim, n_cells=(n_cells,), ref_quad=UnitSquareTensorized(dim=physical_dim, order=quad_order), mapping=Mapping(mappings=[mapping_id]), ) # ── Paramètres DG apprenable ────────────────────────────────────────────────── n_cells_learn = 10 poly_deg_learn = 4 quad_order_learn = 5 N_COLLOC = 6000 N_EPOCHS = 1000 print() print("=" * 60) print("PARTIE 2 — DG apprenable (BasisNN + PINN)") print("=" * 60) nb_basis_learn = poly_deg_learn + 1 mesh_learn = make_mesh(n_cells_learn, quad_order_learn) sigma_learn = (poly_deg_learn + 1) * (poly_deg_learn + physical_dim) / physical_dim class BasisNN(MLP): def __init__(self, key): super().__init__(in_size=1, out_size=1, hidden_sizes=[12, 12], key=key) def __call__(self, y: jnp.ndarray) -> jnp.ndarray: return super().__call__(y) def custom_local_basis(u, coords, i, mesh): val = local_taylor_basis(coords, i, mesh, order=poly_deg_learn, out_dim=out_dim) return val * (2.0 + 0.1 * u[0]) def create_variables(key) -> VariablesDG: pw_basis = PatchwiseParametricBasis( nb_basis=nb_basis_learn, out_dim=out_dim, mesh=mesh_learn, patchwise_parametric_function=BasisNN(key=key), local_basis=custom_local_basis, basis_type="scalar", ) return VariablesDG(basis=pw_basis, nb_variables=out_dim) dx = Segment1D((0.0, 1.0), is_main_domain=True) fClass = make_functional_field_class("f") model = LaplacianDirichletDG(main_domain=dx, f_rhs=fClass(f_rhs), bc="weak") model2 = LaplacianDirichletDG(main_domain=dx, f_rhs=fClass(f_rhs2), bc="weak") key = jax.random.PRNGKey(seed) key, subkey = jax.random.split(key) # pde_learn = LaplacianWeakForm(dim=physical_dim, f=f_rhs) # 1 / n_cells_learn) # assembler_learn = EllipticDGscheme( # pde_learn, # create_variables(key=subkey), # flux_learn, # ) class DGsolverPatchwiseBases(AbstractPhysicNO): flux: SIPGFlux variables: VariablesDG = dynamic() def __init__(self, key, sigma=sigma_learn, h=mesh_learn.h, **newton_kwargs): super().__init__(model_type="x", model_size=1, type_model="scalar") self.flux = SIPGFlux(sigma=sigma, h=h) self.variables = create_variables(key=key) self.newton_kwargs = newton_kwargs def ndof(self) -> int: return self.variables.ndof def alpha_shape(self) -> tuple[int, ...]: return ( self.variables.mesh.n_cells_total, self.variables.trial_basis.nb_basis, self.variables._nb_variables_pre_postprocessing, ) def encoder(self, physical_model: AbstractPhysicalModel) -> NDARRAY_TYPE: assert isinstance(physical_model, LaplacianDirichletDG) assert physical_model.main_domain is not None label = physical_model.main_domain.get_label() residual = physical_model.physical_residuals[label] weakform = LaplacianWeakForm(dim=1, f=residual.f_rhs) # Le schema prend un MODELE, pas une forme faible nue : la condition # de bord vit dans le modele et ne se passe plus au solve. weak_model = AbstractPhysicalWeakModel.from_weak_form( weakform, dirichlet=dirichlet_bc ) assembler = EllipticDGscheme( weak_model, self.variables, self.flux, ) solved = type(assembler).solve( assembler, matrix_free=False, **self.newton_kwargs, ) return solved.variables.dofsl def beta_shape(self) -> tuple[int, ...]: return self.alpha_shape() def propagator(self, alpha: NDARRAY_TYPE) -> NDARRAY_TYPE: # Ici l'encodeur fait DEJA le solve DG et rend les dofs de la solution, # donc le propagateur -- le maillon qui inverse la matrice -- est # l'identite. Un decoupage conforme au schema encodeur / propagateur / # decodeur deplacerait le solve ici ; ce serait un autre exemple. return alpha def decoder(self, alpha: NDARRAY_TYPE, *args: NDARRAY_TYPE) -> NDARRAY_TYPE: var = self.variables return var.local_evaluate_pure(var, alpha, *args) # space = DGEllipticApproximationSpace( # dims={"x": physical_dim, "dofsl": 1}, # list_assemblers=[assembler_learn], # dirichlet_bcs=dirichlet_bc, # model_type="x_dofsl", # newton_kwargs={"max_iter": 1, "tol": 1e-6}, # ) no = DGsolverPatchwiseBases(subkey) space = PhysicNOApproximationSpace( dims={"x": 1}, list_models=[no], model_type="x", ) sampler = TensorizedSampler([DomainSampler(dx)], bc=True, model_type="x") key, sample_dict = sampler.sample(key, N_COLLOC) projector = PhysicNOProjector([model, model2], space, sampler, optimizer="SS-BFGS") key, batched_pdes, _ = projector.sample_physical_models(key, 1) loss_func = projector.build_losses_function() loss_func = jax.jit(loss_func) loss0 = loss_func(space, sample_dict, batched_pdes) print("Loss initiale : ", loss0) print(f"Entraînement ({N_EPOCHS} époques) …") t0 = time.perf_counter() key, projector = projector.project(key, space, N_EPOCHS, 2, N_COLLOC) new_loss = projector.best_loss nspace = projector.space loss_history = projector.losses.losses_history jax.block_until_ready(jax.tree_util.tree_leaves(new_loss)) t_learnable = time.perf_counter() - t0 new_space = projector.space print(f"Loss finale : {new_loss['total']:.6e}") print(f"Temps entraînement (JIT inclus) : {t_learnable:.2f} s") projector.plot([model], exact_sol=u_exact, equal_aspect=False) projector.plot([model2], exact_sol=u_exact2, equal_aspect=False) # # ── Erreurs ───────────────────────────────────────────────────── # fresh_assembler_learn = new_space.assemblers[0] # assembler_solved_learn = EllipticDGscheme.solve( # fresh_assembler_learn, dirichlet_bc, max_iter=1, tol=1e-6 # ) # l2_classical = l2_error( # assembler=assembler_classical, u_exact_fn=u_exact, relative=True # ) # l2_learnable = l2_error( # assembler=assembler_solved_learn, u_exact_fn=u_exact, relative=True # ) # nn = new_space.assemblers[0].variables.trial_basis.patchwise_parametric_function # n_params = sum(p.size for p in jax.tree_util.tree_leaves(nn)) # print() # print( # f"{'':20s} {'Erreur L2':>12s} {'Temps (s)':>10s} {'DOFs':>6s} {'Params NN':>10s}" # ) # print( # f"{'DG classique':20s} {l2_classical:>12.4e} {t_classical:>10.4f} {n_cells_ref * nb_basis_ref:>6d} {'—':>10s}" # ) # print( # f"{'DG apprenable':20s} {l2_learnable:>12.4e} {t_learnable:>10.4f} {n_cells_learn * nb_basis_learn:>6d} {n_params:>10d}" # ) # # ── Résidu PINN ─────────────────────────────────────────────────────────────── # x_plot = jnp.linspace(0.0, 1.0, 300)[:, jnp.newaxis] # u_ref = jax.vmap(u_exact)(x_plot)[:, 0] # u_classical = solved_classical.variables.evaluate(x_plot)[:, 0] # # DG init : solve avec les poids NN initiaux (assembler_learn non entraîné) # assembler_solved_init = EllipticDGscheme.solve( # assembler_learn, dirichlet_bc, max_iter=1, tol=1e-6 # ) # u_learn_init = assembler_solved_init.variables.evaluate(x_plot)[:, 0] # u_learn_final = assembler_solved_learn.variables.evaluate(x_plot)[:, 0] # # Résidu PINN via LaplacianDirichletDG (pinn.evaluator) — PDE-agnostique # dofsl_init = space.get_intermediate_values()[0] # dofsl_final = new_space.get_intermediate_values()[0] # lhs_init = pinn.evaluator.evaluate_physical_residual( # "interior", space, x_plot, dofsl_init # ) # lhs_final = pinn.evaluator.evaluate_physical_residual( # "interior", new_space, x_plot, dofsl_final # ) # rhs_plot = pinn.evaluator.evaluate_rhs_physical_residual("interior", x_plot) # res_init = ((lhs_init - rhs_plot) ** 2)[:, 0] # res_final = ((lhs_final - rhs_plot) ** 2)[:, 0] # # ── Plots ───────────────────────────────────────────────────────────────────── # err_classical = jnp.abs(u_classical - u_ref) # err_learn_final = jnp.abs(u_learn_final - u_ref) # fig, axs = plt.subplots(1, 5, figsize=(25, 5)) # # Panel 1 : solutions # axs[0].plot(x_plot[:, 0], u_ref, "k--", linewidth=2, label="exacte") # axs[0].plot( # x_plot[:, 0], # u_classical, # label=f"DG classique ({n_cells_ref}×{nb_basis_ref} DOFs)", # ) # axs[0].plot(x_plot[:, 0], u_learn_init, ":", alpha=0.5, label="DG appr. (init)") # axs[0].plot( # x_plot[:, 0], # u_learn_final, # label=f"DG apprenable ({n_cells_learn}×{nb_basis_learn} DOFs, {n_params} params)", # ) # axs[0].set_xlabel("x") # axs[0].set_ylabel("u(x)") # axs[0].set_title("Solution DG vs exacte") # axs[0].legend(fontsize=8, loc="upper right") # axs[0].grid(True, alpha=0.3) # # Panel 2 : historique de loss # loss_total = jnp.asarray(pinn.losses.losses_history["total"]).reshape(-1) # axs[1].semilogy(loss_total, "o-", markersize=3) # axs[1].set_xlabel("Époque") # axs[1].set_ylabel("Loss PINN") # axs[1].set_title(f"Historique ({N_EPOCHS} époques)") # axs[1].grid(True, alpha=0.3) # # Panel 3 : temps de calcul (JIT inclus) # labels = [ # f"DG classique\n{t_classical:.4f} s", # f"DG apprenable\n{t_learnable:.2f} s", # ] # axs[2].bar(labels, [t_classical, t_learnable], color=["steelblue", "darkorange"]) # axs[2].set_ylabel("Temps (s, JIT inclus)") # axs[2].set_title("Temps de calcul") # axs[2].set_yscale("log") # axs[2].grid(True, alpha=0.3, axis="y") # # Panel 4 : résidu PINN # axs[3].plot(x_plot[:, 0], res_init, "--", label="initial") # axs[3].plot(x_plot[:, 0], res_final, label="final") # axs[3].axhline(0.0, color="k", linewidth=0.8, linestyle=":") # axs[3].set_xlabel("x") # axs[3].set_ylabel(r"$(-\Delta u_{DG} - f)^2$") # axs[3].set_title("Résidu PINN") # axs[3].legend() # axs[3].grid(True, alpha=0.3) # # Panel 5 : erreur ponctuelle # axs[4].semilogy( # x_plot[:, 0], err_classical, label=f"DG classique (L2 rel={l2_classical:.2e})" # ) # axs[4].semilogy( # x_plot[:, 0], err_learn_final, label=f"DG apprenable (L2 rel={l2_learnable:.2e})" # ) # axs[4].set_xlabel("x") # axs[4].set_ylabel(r"$|u_{DG} - u_{exact}|$") # axs[4].set_title("Erreur ponctuelle") # axs[4].legend(fontsize=8) # axs[4].grid(True, alpha=0.3) # plt.tight_layout() # plt.show()