Metadata-Version: 2.5
Name: yaxlib
Version: 0.1.3
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.svg" 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)`).
- `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/
```
