r"""U-Net physiquement informe sur le laplacien 2D : donnees + Deep Ritz. Meme famille de solutions manufacturees que ``discrete_no/unet_source_to_solution.py``, mais l'operateur est cette fois un **PhysicNO** : =============== ========================================================== encodeur evalue les callables de l'EDP sur la grille (ici la source ; les donnees de BORD sont ecartees, leur second membre recevant la normale en plus du point) propagateur un U-Net : grille -> grille, ce sont les DOF decodeur l'expansion Q1 sur la grille =============== ========================================================== Le decodeur rend la solution evaluable PARTOUT, donc derivable, donc porteuse d'un residu -- ce que la version purement discrete ne permettait pas. **La loss est mixte** : un terme de DONNEES (la solution observee en des points capteurs) et un terme physique de **Deep Ritz**, .. math:: E(u) = \int_\Omega \Big(\tfrac12 |\nabla u|^2 - f u\Big). ⚠ Deux precautions, et elles vont ensemble : - Q1 est C0, donc ses derivees SECONDES sont nulles dans les mailles : un residu fort ``-Delta u = f`` vaudrait identiquement zero et la loss ne verrait rien. Deep Ritz ne demande que des derivees PREMIERES, et c'est ce qui le rend compatible avec cette base. - l'energie doit etre INTEGREE, pas mise au carre : son minimum est negatif, et une MSE minimiserait ``|E|`` donc pousserait vers ``E = 0``. D'ou :class:`ScimbaMean`, qui moyenne l'integrande en gardant son signe. """ import jax import jax.numpy as jnp import matplotlib.pyplot as plt import numpy as np from scimba_jax.domains.meshless_domains.domains_nd import HypercubeND from scimba_jax.neural_operator.data_for_no.grid_data import GridData from scimba_jax.neural_operator.physic_no.grid_based.physic_informed_unet import ( PhysicInformedUnet, ) 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.nonlinear_approximation.optimizers.losses import ScimbaMean, ScimbaMSE from scimba_jax.physical_models.abstract_physical_model import AbstractPhysicalModel from scimba_jax.physical_models.abstract_residuals import InteriorResidual from scimba_jax.physical_models.boundary_residuals import DirichletResidual from scimba_jax.physical_models.data_residuals import CollocDataResidual from scimba_jax.utils.functional_fields import make_functional_field_class jax.config.update("jax_enable_x64", True) N_GRID = 32 # points par direction de la grille des DOF N_MODES = 3 # modes de Fourier de la famille de solutions N_TRAIN, N_TEST = 64, 16 N_SENSORS = 200 # points ou la solution est observee N_EPOCHS_ADAM, N_EPOCHS_BROYDEN, BATCH_SIZE = 1000, 200, 16 N_COLLOC, N_BC_COLLOC = 600, 100 WEIGHT_DATA, WEIGHT_RITZ, WEIGHT_BC = 1.0, 1.0e-2, 1.0 domain = HypercubeND([(0.0, 1.0), (0.0, 1.0)], is_main_domain=True) grid_data = GridData(2, HypercubeND([(0.0, 1.0), (0.0, 1.0)]), (N_GRID, N_GRID)) source_class = make_functional_field_class("laplacian_source") boundary_class = make_functional_field_class("laplacian_boundary") class DeepRitzResidual(InteriorResidual): r"""L'integrande de Deep Ritz, :math:`\tfrac12|\nabla u|^2 - f u`. La source figure DANS l'integrande, pas au second membre : c'est pourquoi la loss associee ignore le second membre. Args: domain: le domaine, f_rhs: la source, model_type: les variables; defaut ``"x"``. """ def __init__(self, domain, f_rhs=None, model_type="x"): super().__init__(domain=domain, size=1, model_type=model_type, f_rhs=f_rhs) def construct_residual(self, *variables): """Construit l'integrande. Args: *variables: les variables de l'espace d'approximation. Returns: l'integrande, comme fonction parametrique. """ u = variables[0] gradient = u.gradient("x") def source_of_x(*args): # La fonction parametrique recoit plus que le point : pour un # espace de NO, les valeurs intermediaires (la sortie du # propagateur) suivent les features. La source, elle, ne depend # que de x -- d'ou cette adaptation, locale au residu qui seul # connait sa propre convention d'arguments. return self.f_rhs(args[0]) return gradient.dot(gradient) * 0.5 - u * source_of_x class LaplacianRitzWithData(AbstractPhysicalModel): """Laplacien 2D : energie de Deep Ritz, bord de Dirichlet, et observations. Args: main_domain: le domaine, f_rhs: la source, f_bc_rhs: la valeur au bord, data: le couple (points capteurs, valeurs observees), model_type: les variables; defaut ``"x"``. """ def __init__(self, main_domain, f_rhs=None, f_bc_rhs=None, data=(), model_type="x"): super().__init__(main_domain=main_domain) self.physical_residuals = { self.main_domain.get_label(): DeepRitzResidual( domain=main_domain, 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, ) self.add_data_residual( "data", CollocDataResidual( size=1, model_type=model_type, data=data, batchable_args=False ), ) # ── La famille de solutions, et les sources qu'on en DEDUIT ────────────────── modes = jnp.arange(1, N_MODES + 1) eigenvalues = jnp.pi**2 * (modes[:, None] ** 2 + modes[None, :] ** 2) def solution_and_source(coefficients): """Rend la solution et sa source, toutes deux exactes. Args: coefficients: les coefficients de la serie, de forme (k, l). Returns: le couple des deux callables. """ def basis(x): return ( jnp.sin(jnp.pi * modes * x[0])[:, None] * jnp.sin(jnp.pi * modes * x[1])[None, :] ) def solution(x): return jnp.sum(coefficients * basis(x))[None] def source(x): return jnp.sum(coefficients * eigenvalues * basis(x))[None] return solution, source def zero_boundary(x, normal): """La condition de bord, nulle. Les seconds membres de bord recoivent la NORMALE en plus du point, d'ou la signature a deux arguments. Args: x: le point du bord, normal: la normale sortante. Returns: zero. """ del x, normal return jnp.zeros(1) def make_batch(key, n_models, sensors): """Tire une famille d'EDP et les observations associees. Args: key: l'etat du generateur, n_models: le nombre de modeles, sensors: les points d'observation. Returns: (nouvelle cle, liste de modeles, valeurs exactes aux capteurs, liste des solutions exactes comme callables). """ key, subkey = jax.random.split(key) coefficients = jax.random.normal(subkey, (n_models, N_MODES, N_MODES)) coefficients = coefficients / (modes[None, :, None] * modes[None, None, :]) models, exact, solutions = [], [], [] for index in range(n_models): solution, source = solution_and_source(coefficients[index]) values = jax.vmap(solution)(sensors) exact.append(values) solutions.append(solution) models.append( LaplacianRitzWithData( main_domain=domain, f_rhs=source_class(source), f_bc_rhs=boundary_class(zero_boundary), data=(sensors, values), ) ) return key, models, jnp.stack(exact), solutions key = jax.random.PRNGKey(0) key, key_sensors = jax.random.split(key) sensors = jax.random.uniform(key_sensors, (N_SENSORS, 2)) key, train_models, train_exact, _ = make_batch(key, N_TRAIN, sensors) key, test_models, test_exact, test_solutions = make_batch(key, N_TEST, sensors) # ── L'operateur, l'espace, le projecteur ───────────────────────────────────── n_fields = PhysicInformedUnet.count_encoded_fields(train_models[0], grid_data) key, key_net = jax.random.split(key) operator = PhysicInformedUnet( grid_data, n_fields, 1, key_net, levels=2, base_channels=8, use_coordinates=True, ) space = PhysicNOApproximationSpace( dims={"x": 2}, list_models=[operator], model_type="x" ) sampler = TensorizedSampler([DomainSampler(domain)], model_type="x", bc=True) interior_label = domain.get_label() losses = { interior_label: (ScimbaMean(WEIGHT_RITZ),), **{boundary: (ScimbaMSE(WEIGHT_BC),) for boundary in train_models[0].boundaries}, "data": (ScimbaMSE(WEIGHT_DATA),), } projector = PhysicNOProjector( train_models, space, sampler, optimizer="Adam", learning_rate=1.0e-3, losses=losses ) print("=" * 74) print("U-Net PHYSIQUEMENT INFORME -- laplacien 2D, donnees + Deep Ritz") print("=" * 74) print(f" grille des DOF : {N_GRID}x{N_GRID} ({N_GRID * N_GRID} Q1 dofs)") print(f" canaux encodes : {n_fields} (les seconds membres d'interieur)") print(f" parametres appris: {operator.ndof()}") print(f" EDP : {N_TRAIN} en entrainement, {N_TEST} en test") print(f" entrainement : {N_EPOCHS_ADAM} Adam puis {N_EPOCHS_BROYDEN} SS-Broyden") key, projector = projector.project( key, space, N_EPOCHS_ADAM, BATCH_SIZE, N_COLLOC, N_BC_COLLOC ) # ── Erreur sur le jeu de TEST ──────────────────────────────────────────────── def relative_errors(trained, models, exact, points): """Erreur L2 relative aux capteurs, modele par modele. Les EDP sont EMPILEES puis evaluees d'un coup, comme dans l'exemple DeepONet : une boucle Python dispatcherait un appel par modele. Args: trained: le projecteur entraine, models: la liste des EDP, exact: les valeurs exactes aux capteurs, points: les points d'observation. Returns: le tableau des erreurs relatives. """ batched = AbstractPhysicalModel.create_batch(models) predicted = jax.vmap(trained.evaluate, in_axes=(0, None))(batched, points) numerator = jnp.linalg.norm(predicted - exact, axis=(1, 2)) denominator = jnp.linalg.norm(exact, axis=(1, 2)) return numerator / denominator def report(trained, label): """Affiche les erreurs d'entrainement et de test. Args: trained: le projecteur entraine, label: le nom de l'etape. Returns: l'erreur de test moyenne. """ on_train = relative_errors(trained, train_models, train_exact, sensors) on_test = relative_errors(trained, test_models, test_exact, sensors) print( f" {label:<12s} entrainement {float(on_train.mean()):.3e}" f" test {float(on_test.mean()):.3e}" ) return float(on_test.mean()) print() error_adam = report(projector, "Adam") # ── Second temps : SS-Broyden REPREND l'espace entraine par Adam ───────────── # Un quasi-Newton part d'ou Adam s'est arrete ; le construire sur un espace # neuf reviendrait a jeter le premier entrainement. trained_space = projector.space projector_broyden = PhysicNOProjector( train_models, trained_space, sampler, optimizer="SS-Broyden", losses=losses ) key, projector_broyden = projector_broyden.project( key, trained_space, N_EPOCHS_BROYDEN, BATCH_SIZE, N_COLLOC, N_BC_COLLOC ) error_broyden = report(projector_broyden, "SS-Broyden") print(f"\n gain du second temps : {error_adam / error_broyden:.2f}x") projector = projector_broyden # les traces montrent le resultat final test_errors = relative_errors(projector, test_models, test_exact, sensors) # ── Traces : exacte, predite, erreur, sur trois EDP de test ────────────────── n_plot = 64 axis_values = jnp.linspace(0.0, 1.0, n_plot) mesh_x, mesh_y = jnp.meshgrid(axis_values, axis_values, indexing="ij") plot_points = jnp.stack([mesh_x.reshape(-1), mesh_y.reshape(-1)], axis=-1) X, Y = np.array(mesh_x), np.array(mesh_y) figure, axes = plt.subplots(3, 3, figsize=(14, 11)) for column in range(3): exact = np.array(jax.vmap(test_solutions[column])(plot_points)[:, 0]).reshape( n_plot, n_plot ) predicted = np.array( projector.evaluate(test_models[column], plot_points)[:, 0] ).reshape(n_plot, n_plot) error = np.abs(predicted - exact) relative = np.linalg.norm(error) / np.linalg.norm(exact) # Echelle partagee sur les DEUX champs : bornee a l'exacte seule, un # depassement de la prediction tomberait hors des niveaux. shared = np.linspace( min(exact.min(), predicted.min()), max(exact.max(), predicted.max()), 21 ) for row, (field, title, cmap) in enumerate( ( (exact, "u exacte", "viridis"), (predicted, "u predite", "viridis"), (error, f"|erreur| (rel. {relative:.1e})", "magma"), ) ): axis = axes[row, column] image = axis.contourf( X, Y, field, levels=shared if row < 2 else 20, cmap=cmap, extend="both" ) figure.colorbar(image, ax=axis, fraction=0.046) axis.set_title(f"{title} (EDP de test {column})", fontsize=10) axis.set_aspect("equal") figure.suptitle( f"U-Net + Q1, donnees + Deep Ritz -- erreur relative de test " f"{float(test_errors.mean()):.2e}" ) plt.tight_layout() plt.show()