r"""U-Net : apprendre l'application source -> solution, avec et sans parametre. Deux cas, tous deux en solutions manufacturees : 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:`UNet` ordinaire apprend :math:`f \mapsto u`. **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)`. Ici 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:`UNetModulated` recoit :math:`\mu = \log_{10}\varepsilon` et module tous ses noyaux par :math:`K(\mu) = W\cdot\varphi(\mu)`. Pour savoir si la modulation sert vraiment, le meme jeu est aussi appris par un :class:`UNet` 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.unet import ( UNet, UNetModulated, ) 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 BASE_CHANNELS = 4 # petit reseau, volontairement N_EPOCHS, BATCH_SIZE, LEARNING_RATE = 1500, 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 (UNet ordinaire)") print("=" * 74) unet = UNet( grid_2d, 1, 1, key_net, levels=3, base_channels=BASE_CHANNELS, use_coordinates=True ) print(f" grille {N_2D}x{N_2D}, {unet.ndof()} parametres, {N_TRAIN} echantillons") projector_2d = train(unet, (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 = UNetModulated( grid_1d, 1, 1, key_mod, levels=3, base_channels=BASE_CHANNELS, param_size=1, latent_size=8, ) # ⚠ Le temoin est ELARGI a dessein. La modulation multiplie chaque noyau par # latent_size, donc a base_channels egal le modulé aurait ~8 fois plus de # poids : sa victoire ne dirait alors rien sur la modulation, seulement sur la # taille. base_channels=12 lui donne un budget comparable (91 729 contre # 87 929), et la seule difference qui reste est que l'un voit eps et l'autre # non. plain = UNet(grid_1d, 1, 1, key_plain, levels=3, base_channels=12) 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" UNet MODULE : {train_mod:.3e} {error_mod:.3e}" f" {error_mod / train_mod:5.1f}x" ) print( f" UNet 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), donc tout un reseau est decrit par\n" " latent_size coefficients de melange.\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." ) # ══════════════════════════════════════════════════════════════════════════════ # 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. L'echelle est prise sur les DEUX # champs : bornee a l'exacte seule, tout depassement de la prediction # tomberait hors des niveaux et `contourf` laisserait un trou blanc. 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"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="UNet module") top.plot(abscissa, plain_prediction, ":", lw=1.8, label="UNet 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"Advection-diffusion 1D — erreur relative moyenne sur {N_TEST} cas de test : " f"module {error_mod:.2e} temoin {error_plain:.2e}" ) plt.tight_layout() plt.show()