r"""FNO : apprendre l'application source -> solution, avec et sans parametre. Exactement les DEUX MEMES cas que ``unet_source_to_solution.py`` -- memes solutions manufacturees, memes tirages, memes tailles de jeux -- pour que les deux architectures soient comparables ligne a ligne. On se donne la solution :math:`u`, on en DEDUIT la source analytiquement, et on apprend la correspondance inverse :math:`f \mapsto u`. Aucune difference finie n'intervient, donc aucune erreur de discretisation ne vient polluer la mesure. **1. Laplacien 2D, sans parametre.** Sur :math:`(0,1)^2`, .. math:: u(x,y) = \sum_{k,l} a_{kl}\,\sin(k\pi x)\sin(l\pi y), \qquad -\Delta u = \sum_{k,l} a_{kl}\,(k^2+l^2)\pi^2 \sin(k\pi x)\sin(l\pi y). Un :class:`FNO` ordinaire apprend :math:`f \mapsto u`. C'est le cas le plus favorable qui soit pour lui : la famille de solutions EST engendree par les premiers modes de Fourier, et l'operateur inverse y est exactement diagonal, de symbole :math:`1/((k^2+l^2)\pi^2)`. Le FNO n'a donc, en principe, qu'a apprendre ce symbole -- ce qu'un U-Net doit reconstituer par des convolutions locales empilees. **2. Advection-diffusion 1D, parametree par eps.** Sur :math:`(0,1)`, .. math:: -\varepsilon u'' + u' = f, \qquad u(x) = \sum_k a_k \sin(k\pi x), d'ou :math:`f = \varepsilon\sum_k a_k (k\pi)^2 \sin(k\pi x) + \sum_k a_k\,k\pi\cos(k\pi x)`. L'operateur a inverser CHANGE avec :math:`\varepsilon` : a :math:`\varepsilon` grand la diffusion domine, a :math:`\varepsilon` petit c'est le transport. Un :class:`FNOModulated` recoit :math:`\mu = \log_{10}\varepsilon` et module :math:`W\cdot\varphi(\mu)` PARTOUT -- relevement, ponderation des modes, melange de canaux, projection. Pour savoir si la modulation sert vraiment, le meme jeu est aussi appris par un :class:`FNO` ordinaire, qui ne voit pas :math:`\varepsilon` : c'est le temoin. Il est ELARGI pour avoir un budget de parametres comparable, sans quoi on mesurerait la taille des reseaux et non l'apport du parametre. L'exemple affiche les erreurs d'ENTRAINEMENT et de TEST pour les deux : c'est leur ecart qui distingue un reseau qui manque de capacite d'un reseau qui sur-apprend, et les deux se ressemblent si l'on ne regarde que le test. """ 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.discrete_no.grid_based.fno import ( FNO, FNOModulated, ) from scimba_jax.nonlinear_approximation.numerical_solvers.no_projectors import ( NOProjector, ) jax.config.update("jax_enable_x64", True) N_MODES = 4 # modes de Fourier de la famille de solutions N_TRAIN, N_TEST = 256, 64 HIDDEN_CHANNELS = 8 # petit reseau, volontairement FNO_MODES = 8 # modes RETENUS par le reseau -- deux fois ceux du signal N_BLOCKS = 3 N_EPOCHS, BATCH_SIZE, LEARNING_RATE = 1200, 32, 3.0e-3 def relative_error(prediction: jnp.ndarray, target: jnp.ndarray) -> float: """Erreur L2 relative sur tout le jeu. Args: prediction: les valeurs predites, target: les valeurs visees. Returns: l'erreur relative. """ return float(jnp.linalg.norm(prediction - target) / jnp.linalg.norm(target)) def train(operator, data, key, label): """Entraine un operateur et rend le projecteur entraine. Args: operator: l'operateur a entrainer, data: le jeu de donnees passe a NOProjector, key: l'etat du generateur, label: le nom affiche dans la barre de progression. Returns: le projecteur entraine. """ projector = NOProjector(operator, data, learning_rate=LEARNING_RATE) _, projector = projector.project( key, operator, N_EPOCHS, batch_size=BATCH_SIZE, tqdm_desc=label ) return projector # ══════════════════════════════════════════════════════════════════════════════ # 1. Laplacien 2D : -Delta u = f, sans parametre # ══════════════════════════════════════════════════════════════════════════════ N_2D = 64 # grille fine grid_2d = GridData(2, HypercubeND([(0.0, 1.0), (0.0, 1.0)]), (N_2D, N_2D)) x_2d, y_2d = grid_2d.grid[..., 0], grid_2d.grid[..., 1] modes = jnp.arange(1, N_MODES + 1) sin_x = jnp.sin(jnp.pi * modes[:, None, None] * x_2d[None]) # (k, n, n) sin_y = jnp.sin(jnp.pi * modes[:, None, None] * y_2d[None]) # (l, n, n) eigenvalues = jnp.pi**2 * (modes[:, None] ** 2 + modes[None, :] ** 2) # (k, l) def laplacian_dataset(key, n_samples): """Tire des solutions et en deduit les sources exactes. Args: key: l'etat du generateur, n_samples: le nombre d'echantillons. Returns: le couple (sources, solutions), de forme (n, N, N, 1). """ coefficients = jax.random.normal(key, (n_samples, N_MODES, N_MODES)) coefficients = coefficients / (modes[None, :, None] * modes[None, None, :]) basis = sin_x[:, None] * sin_y[None, :] # (k, l, n, n) solution = jnp.einsum("skl,klxy->sxy", coefficients, basis) source = jnp.einsum("skl,kl,klxy->sxy", coefficients, eigenvalues, basis) return source[..., None], solution[..., None] key = jax.random.PRNGKey(0) key, key_train, key_test, key_net, key_fit = jax.random.split(key, 5) f_train, u_train = laplacian_dataset(key_train, N_TRAIN) f_test, u_test = laplacian_dataset(key_test, N_TEST) # Une seule constante, lue sur l'entrainement : les sources sont ~100 fois plus # grandes que les solutions (les valeurs propres du laplacien). scale = float(jnp.std(f_train)) print("\n" + "=" * 74) print("1. LAPLACIEN 2D -Delta u = f (FNO ordinaire)") print("=" * 74) fno = FNO( grid_2d, 1, 1, key_net, n_modes=FNO_MODES, hidden_channels=HIDDEN_CHANNELS, n_blocks=N_BLOCKS, use_coordinates=True, ) print(f" grille {N_2D}x{N_2D}, {fno.ndof()} parametres, {N_TRAIN} echantillons") projector_2d = train(fno, (f_train / scale, u_train), key_fit, "laplacien 2D") prediction_2d = projector_2d.evaluate(f_test / scale) print(f"\n loss finale : {float(projector_2d.best_loss['total']):.3e}") print(f" erreur relative (test) : {relative_error(prediction_2d, u_test):.3e}") # ══════════════════════════════════════════════════════════════════════════════ # 2. Advection-diffusion 1D : -eps u'' + u' = f, eps est LE parametre # ══════════════════════════════════════════════════════════════════════════════ N_1D = 64 grid_1d = GridData(1, HypercubeND([(0.0, 1.0)]), (N_1D,)) x_1d = grid_1d.grid[..., 0] sin_1d = jnp.sin(jnp.pi * modes[:, None] * x_1d[None]) # (k, n) cos_1d = jnp.cos(jnp.pi * modes[:, None] * x_1d[None]) # (k, n) EPS_MIN, EPS_MAX = 1.0e-2, 1.0 def advection_diffusion_dataset(key, n_samples): """Tire (solution, eps) et en deduit la source exacte. Args: key: l'etat du generateur, n_samples: le nombre d'echantillons. Returns: le triplet (sources, parametres, solutions). """ key_a, key_eps = jax.random.split(key) coefficients = jax.random.normal(key_a, (n_samples, N_MODES)) / modes[None, :] # eps tire log-uniformement : deux decades, donc les deux regimes. log_eps = jax.random.uniform( key_eps, (n_samples,), minval=jnp.log10(EPS_MIN), maxval=jnp.log10(EPS_MAX) ) eps = 10.0**log_eps solution = jnp.einsum("sk,kx->sx", coefficients, sin_1d) diffusion = jnp.einsum("sk,k,kx->sx", coefficients, (jnp.pi * modes) ** 2, sin_1d) advection = jnp.einsum("sk,k,kx->sx", coefficients, jnp.pi * modes, cos_1d) source = eps[:, None] * diffusion + advection return source[..., None], log_eps[:, None], solution[..., None] key, key_train, key_test, key_mod, key_plain, key_fit = jax.random.split(key, 6) f1_train, mu_train, u1_train = advection_diffusion_dataset(key_train, N_TRAIN) f1_test, mu_test, u1_test = advection_diffusion_dataset(key_test, N_TEST) scale_1d = float(jnp.std(f1_train)) print("\n" + "=" * 74) print("2. ADVECTION-DIFFUSION 1D -eps u'' + u' = f (eps = parametre)") print("=" * 74) modulated = FNOModulated( grid_1d, 1, 1, key_mod, n_modes=FNO_MODES, hidden_channels=HIDDEN_CHANNELS, n_blocks=N_BLOCKS, param_size=1, latent_size=8, # Pas de coordonnees : elles viennent de la grille declaree et fixeraient # la resolution, ce qui interdirait la section 3. use_coordinates=False, ) # ⚠ Le temoin est ELARGI a dessein. La modulation multiplie chaque tenseur de # modes par latent_size, donc a hidden_channels egal le modulé aurait ~8 fois # plus de poids : sa victoire ne dirait alors rien sur la modulation, seulement # sur la taille. On elargit le temoin jusqu'a un budget comparable, et la seule # difference qui reste est que l'un voit eps et l'autre non. plain = FNO( grid_1d, 1, 1, key_plain, n_modes=FNO_MODES, hidden_channels=3 * HIDDEN_CHANNELS, n_blocks=N_BLOCKS, use_coordinates=False, ) print(f" module : {modulated.ndof():6d} parametres (voit eps)") print(f" temoin : {plain.ndof():6d} parametres (ne voit PAS eps, budget comparable)") projector_mod = train( modulated, (f1_train / scale_1d, mu_train, u1_train), key_fit, "module" ) projector_plain = train(plain, (f1_train / scale_1d, u1_train), key_fit, "temoin") # Erreurs sur l'entrainement ET sur le test : c'est leur ECART qui dit ce qui # se passe, pas le test seul. train_mod = relative_error( projector_mod.evaluate(f1_train / scale_1d, mu_train), u1_train ) train_plain = relative_error(projector_plain.evaluate(f1_train / scale_1d), u1_train) error_mod = relative_error(projector_mod.evaluate(f1_test / scale_1d, mu_test), u1_test) error_plain = relative_error(projector_plain.evaluate(f1_test / scale_1d), u1_test) print("\n entrainement test test/entrainement") print( f" FNO MODULE : {train_mod:.3e} {error_mod:.3e}" f" {error_mod / train_mod:5.1f}x" ) print( f" FNO temoin : {train_plain:.3e} {error_plain:.3e}" f" {error_plain / train_plain:5.1f}x" ) print( f"\n rapport des erreurs de TEST, temoin / module : {error_plain / error_mod:.1f}x" ) print( "\n => Regarder la derniere colonne avant de conclure. Si le temoin ajuste\n" " l'entrainement AUSSI BIEN ou MIEUX que le module tout en echouant sur le\n" " test, alors il ne manque pas de capacite : il SUR-APPREND. La modulation\n" " agit alors comme une regularisation -- les poids du module ne sont pas\n" " libres, ils sont lies par phi(mu).\n" " A noter : eps est en principe LISIBLE dans f, les modes en sinus portant\n" " eps*a_k*(k.pi)^2 et ceux en cosinus a_k*k.pi, orthogonaux entre eux. La\n" " tache du temoin n'est donc pas ambigue, seulement plus dure." ) # ══════════════════════════════════════════════════════════════════════════════ # 3. Ce qu'un U-Net ne sait pas faire : changer de resolution # ══════════════════════════════════════════════════════════════════════════════ # Les poids sont indexes par MODE, pas par point de grille : le meme reseau # repond sur une grille deux fois plus fine, SANS reentrainement. On le verifie # sur le cas 1D, en re-echantillonnant la meme famille de solutions. N_FINE = 128 x_fine = jnp.linspace(0.0, 1.0, N_FINE) sin_fine = jnp.sin(jnp.pi * modes[:, None] * x_fine[None]) cos_fine = jnp.cos(jnp.pi * modes[:, None] * x_fine[None]) key, key_fine = jax.random.split(key) key_a, key_eps = jax.random.split(key_fine) coefficients = jax.random.normal(key_a, (N_TEST, N_MODES)) / modes[None, :] log_eps_fine = jax.random.uniform( key_eps, (N_TEST,), minval=jnp.log10(EPS_MIN), maxval=jnp.log10(EPS_MAX) ) eps_fine = 10.0**log_eps_fine u_fine = jnp.einsum("sk,kx->sx", coefficients, sin_fine)[..., None] f_fine = ( eps_fine[:, None] * jnp.einsum("sk,k,kx->sx", coefficients, (jnp.pi * modes) ** 2, sin_fine) + jnp.einsum("sk,k,kx->sx", coefficients, jnp.pi * modes, cos_fine) )[..., None] # Le modele ENTRAINE vit dans le projecteur, pas dans "modulated" : c'est # projector_mod.operator (ou .evaluate(), utilise partout ailleurs). prediction_fine = projector_mod.evaluate(f_fine / scale_1d, log_eps_fine[:, None]) print("\n" + "=" * 74) print("3. MEMES POIDS, GRILLE DEUX FOIS PLUS FINE (aucun reentrainement)") print("=" * 74) print(f" entraine sur {N_1D} points, evalue sur {N_FINE} :") print(f" erreur relative : {relative_error(prediction_fine, u_fine):.3e}") print( " => a comparer a l'erreur de test ci-dessus. Un U-Net ne peut meme pas\n" " etre appele ici : ses poids sont indexes par pixel." ) # ══════════════════════════════════════════════════════════════════════════════ # Traces # ══════════════════════════════════════════════════════════════════════════════ # --- 2D : solution exacte, prediction, erreur, sur trois echantillons -------- samples_2d = (0, 1, 2) figure_2d, axes_2d = plt.subplots( 3, len(samples_2d), figsize=(4.5 * len(samples_2d), 11) ) X, Y = np.array(x_2d), np.array(y_2d) for column, sample in enumerate(samples_2d): exact = np.array(u_test[sample, ..., 0]) predicted = np.array(prediction_2d[sample, ..., 0]) error = np.abs(predicted - exact) relative = np.linalg.norm(error) / np.linalg.norm(exact) # Exacte et predite partagent leur echelle, sinon l'oeil compare deux # colormaps differentes et ne voit rien. low = min(exact.min(), predicted.min()) high = max(exact.max(), predicted.max()) shared_levels = np.linspace(low, high, 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_2d[row, column] levels = shared_levels if row < 2 else 20 image = axis.contourf(X, Y, field, levels=levels, cmap=cmap, extend="both") figure_2d.colorbar(image, ax=axis, fraction=0.046) axis.set_title(f"{title} (echantillon {sample})", fontsize=10) axis.set_aspect("equal") figure_2d.suptitle( f"FNO — Laplacien 2D — erreur relative moyenne sur {N_TEST} cas de test : " f"{relative_error(prediction_2d, u_test):.2e}" ) plt.tight_layout() plt.show() # --- 1D : les trois courbes, puis les erreurs ponctuelles -------------------- prediction_mod = projector_mod.evaluate(f1_test / scale_1d, mu_test) prediction_plain = projector_plain.evaluate(f1_test / scale_1d) order = np.argsort(np.array(mu_test[:, 0])) samples_1d = (order[0], order[len(order) // 2], order[-1]) # eps petit, moyen, grand figure_1d, axes_1d = plt.subplots(2, len(samples_1d), figsize=(5 * len(samples_1d), 8)) abscissa = np.array(x_1d) for column, sample in enumerate(samples_1d): eps_value = 10.0 ** float(mu_test[sample, 0]) exact = np.array(u1_test[sample, :, 0]) modulated_prediction = np.array(prediction_mod[sample, :, 0]) plain_prediction = np.array(prediction_plain[sample, :, 0]) top = axes_1d[0, column] top.plot(abscissa, exact, "k", lw=2, label="exacte") top.plot(abscissa, modulated_prediction, "--", lw=1.8, label="FNO module") top.plot(abscissa, plain_prediction, ":", lw=1.8, label="FNO temoin") top.set_title(rf"$\varepsilon$ = {eps_value:.3f}") top.legend(fontsize=8) top.grid(alpha=0.3) bottom = axes_1d[1, column] bottom.semilogy( abscissa, np.abs(modulated_prediction - exact) + 1e-16, "--", label="module" ) bottom.semilogy( abscissa, np.abs(plain_prediction - exact) + 1e-16, ":", label="temoin" ) bottom.set_title("|erreur| ponctuelle", fontsize=10) bottom.set_xlabel("x") bottom.legend(fontsize=8) bottom.grid(alpha=0.3, which="both") figure_1d.suptitle( f"FNO — Advection-diffusion 1D — erreur relative moyenne sur {N_TEST} cas " f"de test : module {error_mod:.2e} temoin {error_plain:.2e}" ) plt.tight_layout() plt.show()