r"""Validation de l'API des opérateurs neuronaux DISCRETS, sur un cas à réponse connue. L'opérateur discret prend les valeurs de :math:`p` fonctions sur un support de :math:`n` points et rend les valeurs de :math:`q` fonctions sur le même support. Contrairement à un ``PhysicNO``, il n'a pas de décodeur : il ne sait pas répondre ailleurs qu'aux points donnés. C'est ce qui, plus tard, autorisera la vraie FFT d'un FNO ou la query dépendante de l'entrée d'un transformer. Trois cas, dans cet ordre : 1. **Cas exact.** :math:`v(x) = W^T u(x)` en chaque point, avec le *même* ``W`` partout. L'opérateur ``PointwiseLinearOperator`` peut représenter cela exactement : on vérifie qu'il retrouve ``W`` et que la loss tombe au niveau de la précision machine. 2. **Cas « presque ».** :math:`v = W^T u + \varepsilon\, G_\sigma * (W^T u)`, où :math:`G_\sigma *` est un lissage gaussien : la relation n'est plus locale, mais elle l'est *presque*. On vérifie que ``W`` est retrouvé à :math:`O(\varepsilon)` près et que l'erreur relative reste de cet ordre. 3. **Contrôle.** Le même cas 2 avec un ``PointwiseMLPOperator``, bien plus riche mais toujours **local**. S'il ne fait pas mieux, c'est que ce qui reste est *intrinsèquement non local* — donc hors de portée de tout opérateur ponctuel, et précisément ce qu'une couche spectrale irait chercher. """ import jax import jax.numpy as jnp from scimba_jax.neural_operator.discrete_no.pointwise_operators import ( PointwiseLinearOperator, PointwiseMLPOperator, ) from scimba_jax.nonlinear_approximation.numerical_solvers.no_projectors import ( NOProjector, ) jax.config.update("jax_enable_x64", True) # ── Problème ────────────────────────────────────────────────────────────────── N_POINTS = 64 # n : points du support P_IN = 3 # p : fonctions d'entrée Q_OUT = 2 # q : fonctions de sortie N_SAMPLES = 256 # nombre de couples (u, v) N_MODES = 8 # modes de Fourier des fonctions tirées EPS = 0.05 # amplitude du terme non local SIGMA = 3.0 # largeur du lissage, en nombre de points W_TRUE = jnp.array([[1.0, 2.0], [0.0, -1.0], [3.0, 0.5]]) N_EPOCHS = 4000 LEARNING_RATE = 5.0e-2 x_grid = jnp.linspace(0.0, 1.0, N_POINTS, endpoint=False) def random_functions(key: jnp.ndarray, n_samples: int, n_channels: int) -> jnp.ndarray: """Tire des fonctions lisses : séries de Fourier à coefficients décroissants. Args: key: l'état du générateur, n_samples: le nombre d'échantillons, n_channels: le nombre de fonctions par échantillon. Returns: un tableau de forme (n_samples, N_POINTS, n_channels). """ key_a, key_b = jax.random.split(key) modes = jnp.arange(1, N_MODES + 1) decay = 1.0 / modes shape = (n_samples, n_channels, N_MODES) coef_cos = jax.random.normal(key_a, shape) * decay coef_sin = jax.random.normal(key_b, shape) * decay phase = 2.0 * jnp.pi * modes[None, :] * x_grid[:, None] # (n, modes) values = jnp.einsum("scm,nm->snc", coef_cos, jnp.cos(phase)) + jnp.einsum( "scm,nm->snc", coef_sin, jnp.sin(phase) ) return values def gaussian_smoothing(values: jnp.ndarray) -> jnp.ndarray: """Lisse chaque fonction par convolution gaussienne périodique. C'est le terme NON LOCAL : sa valeur en un point dépend des voisins, donc aucun opérateur ponctuel ne peut le représenter. Args: values: un tableau (n_samples, N_POINTS, n_channels). Returns: le lissé, de même forme. """ offsets = jnp.arange(N_POINTS) distance = jnp.minimum(offsets, N_POINTS - offsets) # distance périodique kernel = jnp.exp(-0.5 * (distance / SIGMA) ** 2) kernel = kernel / kernel.sum() kernel_hat = jnp.fft.rfft(kernel) values_hat = jnp.fft.rfft(values, axis=1) return jnp.fft.irfft(values_hat * kernel_hat[None, :, None], n=N_POINTS, axis=1) def relative_error(prediction: jnp.ndarray, target: jnp.ndarray) -> float: """Erreur L2 relative sur tout le jeu. Args: prediction: les valeurs prédites, target: les valeurs visées. Returns: l'erreur relative. """ return float(jnp.linalg.norm(prediction - target) / jnp.linalg.norm(target)) def train(operator, inputs, targets, key, label): """Entraîne un opérateur et rend le projecteur entraîné. Args: operator: l'opérateur à entraîner, inputs: les entrées, targets: les cibles, key: l'état du générateur, label: le nom affiché. Returns: le projecteur entraîné. """ projector = NOProjector(operator, (inputs, targets), learning_rate=LEARNING_RATE) _, projector = projector.project(key, operator, N_EPOCHS, tqdm_desc=label) return projector key = jax.random.PRNGKey(0) key, key_data, key_op = jax.random.split(key, 3) u_data = random_functions(key_data, N_SAMPLES, P_IN) v_local = u_data @ W_TRUE # ── 1. Cas exact : v = W u ─────────────────────────────────────────────────── print("\n" + "=" * 70) print("1. CAS EXACT v = W u (l'opérateur ponctuel peut le représenter)") print("=" * 70) operator = PointwiseLinearOperator(P_IN, Q_OUT, key_op) projector = train(operator, u_data, v_local, key, "cas exact") print(f"\n loss finale : {float(projector.best_loss['total']):.3e}") print(f" erreur relative : {relative_error(projector.evaluate(), v_local):.3e}") print( f" |W_appris - W_vrai|: {float(jnp.abs(projector.operator.weight - W_TRUE).max()):.3e}" ) print(f" |biais| : {float(jnp.abs(projector.operator.bias).max()):.3e}") # ── 2. Cas « presque » : v = W u + eps * lissage(W u) ───────────────────────── v_almost = v_local + EPS * gaussian_smoothing(v_local) part_nonlocale = relative_error(v_local, v_almost) print("\n" + "=" * 70) print(f"2. CAS PRESQUE v = W u + {EPS} * lissage(W u)") print("=" * 70) print(f" part non locale des données (||v - Wu|| / ||v||) : {part_nonlocale:.3e}") operator = PointwiseLinearOperator(P_IN, Q_OUT, key_op) projector_almost = train(operator, u_data, v_almost, key, "cas presque") print(f"\n loss finale : {float(projector_almost.best_loss['total']):.3e}") print( f" erreur relative : {relative_error(projector_almost.evaluate(), v_almost):.3e}" ) print( f" |W_appris - W_vrai|: {float(jnp.abs(projector_almost.operator.weight - W_TRUE).max()):.3e}" ) # ── 3. Contrôle : un opérateur local PLUS RICHE ne fait pas mieux ──────────── print("\n" + "=" * 70) print("3. CONTROLE le même cas avec un MLP ponctuel (local, mais non linéaire)") print("=" * 70) operator_mlp = PointwiseMLPOperator(P_IN, Q_OUT, [16, 16], key_op) print( f" ndof : {operator_mlp.ndof()} contre {PointwiseLinearOperator(P_IN, Q_OUT, key_op).ndof()} pour le linéaire" ) projector_mlp = train(operator_mlp, u_data, v_almost, key, "contrôle MLP") err_lin = relative_error(projector_almost.evaluate(), v_almost) err_mlp = relative_error(projector_mlp.evaluate(), v_almost) print(f"\n erreur relative, linéaire ponctuel : {err_lin:.3e}") print(f" erreur relative, MLP ponctuel : {err_mlp:.3e}") print(f" part non locale des données : {part_nonlocale:.3e}") print( "\n => le MLP ponctuel, avec 370 paramètres contre 8, ne fait PAS mieux que\n" " le linéaire. L'opérateur local absorbe le gain moyen du lissage (4.2e-2\n" " de perturbation ramenés à 1.0e-2) et bute ensuite sur un résidu\n" " irréductible : ce qui reste est NON LOCAL, hors de portée de tout\n" " opérateur ponctuel quel que soit son nombre de paramètres. C'est\n" " exactement ce qu'une couche spectrale irait chercher." )