Metadata-Version: 2.5
Name: yaxlib
Version: 0.1.6
Summary: Réseaux de neurones au plus près du jax de base : modules-pytrees stricts, zéro machinerie de filtrage (import : yax)
Author: Vincent Vigon
License: MIT
Requires-Python: >=3.10
Requires-Dist: jax>=0.4.30
Requires-Dist: matplotlib
Requires-Dist: numpy
Requires-Dist: optax>=0.2
Provides-Extra: dev
Requires-Dist: pytest; extra == 'dev'
Description-Content-Type: text/markdown

<p align="center">
  <!-- URL absolue : necessaire pour l'affichage sur PyPI -->
  <img src="https://octaviogame.com/recherche/yax/logo_complex.png" width="340" alt="YA+ — le logo de yaxlib : un yack, un A et un +, imbriques">
</p>

# yaxlib

Un mini-framework de réseaux de neurones pour jax.
[Documentation](https://octaviogame.com/recherche/yax/) —
distribution `yaxlib`, import `yax` :

```bash
pip install yaxlib          # ou : pip install <url du wheel>
```

```python
import jax.random as jr
import yax

model = yax.MLP((2, 32, 32, 1), "tanh", jr.key(0))
```

## Pourquoi yax ?

L'atout de yax est la **simplicité** : rester au plus près du jax de base. La
hiérarchie des frameworks se mesure en concepts additionnels — flax introduit
ses collections de variables, ses scopes et son cycle init/apply ; equinox
réduit cela à des modules-pytrees, mais y ajoute sa machinerie de filtrage
(`filter_grad`, `partition`/`combine`) et son drapeau `inference` dans les
feuilles. yax n'ajoute que deux idées : le **module-pytree strict** (les
feuilles sont exactement les paramètres, tout le reste est statique) et la
**signature `apply(x, rkey)`**. Conséquence : `jax.grad`, `jax.jit`,
`jax.vmap` et optax s'utilisent *nus*, exactement comme dans la documentation
jax — rien à désapprendre, rien à envelopper.

Né pour un cours, yax est dimensionné pour servir au-delà : des modèles de
recherche compacts, lisibles, et un périmètre volontairement réduit — ce qui
n'y est pas se code en jax ordinaire, sans friction.

## Principes

- **Un modèle est un pytree.** `yax.Module` range les tableaux dans les
  feuilles et tout le reste (`yax.StaticField`) dans la structure :
  `jax.grad(loss)(model)`, `jax.jit` et `optimizer.init(model)` acceptent le
  modèle tel quel, sans machinerie de filtrage. Les champs dynamiques ne
  peuvent contenir que des tableaux jax, des sous-modules ou des conteneurs de
  ceux-ci — tout écart est une erreur immédiate et explicite à la
  construction. Un `StaticField` peut contenir un tableau : il devient une
  constante du modèle (encodage positionnel, grille figée), invisible pour les
  gradients.
- **Signature uniforme `apply(x, rkey=None)`**, écrite pour UN échantillon
  (le batch vient de `jax.vmap`). `rkey` est une *source d'aléatoire*
  (dropout, échantillonnage), jamais un mode.
- **Le mode se bascule par `model = model.set_inference(True/False)`**
  (récursif, immuable). Le `Trainer` inclus dans `yax` entraîne
  en `False`, valide et rend le meilleur modèle en `True`. Mais `yax` peut aussi s'utiliser sans ce `Trainer`. 
- **Immutabilité** : on « modifie » un module avec `yax.tree_at`.

## Contenu

- `yax.layers` : Linear, MLP, Dropout, PReLU, LayerNorm, Embedding, Conv_nd
  (convolution 1D/2D/3D...), RNN_layer (GRU/LSTM), MultiHeadAttention,
  TransformerBlock, MessagePassing_layer, encodage positionnel.
- `yax.models` : UNet_nd et FNO_nd (opérateur neuronal de Fourier — le même
  modèle s'évalue à n'importe quelle résolution), tous deux en dimension
  quelconque ; MiniYOLO ; et trois modèles génératifs — VAE, RealNVP
  (flot normalisant à vraisemblance exacte), Diffusion (DDPM) — compacts et
  lisibles, chacun avec sa perte et sa méthode d'échantillonnage.
- `yax.training` : Trainer (checkpoints par `mother_folder`), History, pertes
  (`loss_fn(model, x, y, rkey)`), samplers. Le Trainer ne connaît pas les
  données : une époque = un appel au *sampler* puis une validation.
  `DatasetSampler(X, Y, batch_size)` pour un jeu fini (mélangé sans remise),
  `FunctionSampler(f, nb_batches)` pour des batchs générés — physique, PINN,
  Ritz, où `y` vaut `None`. La validation est un batch fixe `(x_val, y_val)`
  ou un sampler appelé avec une clé constante (même jeu à toutes les époques).
  L'optimiseur est dans la config, sous forme de données : `optimizer` est un nom
  du dictionnaire `yax.OPTIMIZERS` (`"adam"` par défaut, `"adamw"`, `"lion"`,
  `"sgd"`...) et `optimizer_options` ses réglages —
  `TrainConfig(5e-3, 300, optimizer="adamw", optimizer_options={"weight_decay": 0.1})`.
  La config étant enregistrée avec le run, celui-ci dit à lui seul comment il a
  été entraîné. Le Trainer fabrique le schedule (seul à connaître le nombre de
  pas) et le passe au constructeur. `optimizer="lbfgs"` marche aussi : le
  Trainer fournit à tous les optimiseurs la valeur de la perte et la fonction
  qui la calcule — ce que réclame une recherche linéaire, et que les autres
  ignorent. Rien à réécrire côté utilisateur ; il faut en revanche un seul
  batch, fixe (`DatasetSampler(X, Y, len(X))`). Enfin `patience` arrête la
  boucle après N époques sans record de validation — un budget épargné, pas un
  gain de qualité : le modèle rendu est de toute façon le meilleur.
- `yax.image` : augmentation différentiable et vmap-able.
- `demos/` : une démonstration synthétique par famille de modèles, qui
  converge en quelques secondes sur CPU.

## Tests

```bash
pip install yaxlib[dev]     # ajoute pytest
pytest tests/
```
