Metadata-Version: 2.5
Name: yaxlib
Version: 0.1.1
Summary: Un mini-framework de réseaux de neurones pour jax, à but pédagogique (import : yax)
Author: Vincent Vigon
License: MIT
Requires-Python: >=3.10
Requires-Dist: jax>=0.4.30
Requires-Dist: numpy
Requires-Dist: optax>=0.2
Provides-Extra: dev
Requires-Dist: equinox; extra == 'dev'
Requires-Dist: matplotlib; extra == 'dev'
Requires-Dist: pytest; extra == 'dev'
Description-Content-Type: text/markdown

# yaxlib

Un mini-framework de réseaux de neurones pour jax, à but pédagogique.
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))
```

## 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). En pratique on n'y touche pas : le `Trainer` entraîne
  en `False`, valide et rend le meilleur modèle en `True`.
- **Immutabilité** : on « modifie » un module avec `yax.tree_at`.

## Contenu

- `yax.layers` : Linear, MLP, Dropout, LayerNorm, Embedding, Conv_layer,
  RNN_layer (GRU/LSTM), MultiHeadAttention, TransformerBlock,
  MessagePassing_layer, encodage positionnel.
- `yax.models` : UNet, MiniYOLO (références des mini-projets).
- `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]
pytest tests/
```

Les couches réimplémentées (GRU, LSTM, convolution) sont vérifiées
numériquement contre `equinox.nn`, qui ne sert qu'à cela.
