Metadata-Version: 2.5
Name: yaxlib
Version: 0.2.2
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="300" alt="yax">
</p>

# yax — des réseaux de neurones au plus près de jax

[![PyPI](https://img.shields.io/pypi/v/yaxlib?label=yaxlib)](https://pypi.org/project/yaxlib/)
[![Python](https://img.shields.io/badge/python-3.10%2B-blue)](https://pypi.org/project/yaxlib/)
[![Documentation](https://img.shields.io/badge/documentation-octaviogame.com-4051b5)](https://octaviogame.com/recherche/yax/)

**yax** sert à écrire, entraîner et retrouver des réseaux de neurones avec jax,
sans rien apprendre d'autre que jax. Un modèle est une classe Python dont les
tableaux sont les paramètres : `jax.grad`, `jax.jit`, `jax.vmap` et optax
s'appliquent dessus directement. Une fonction, `yax.training.train`, fait la
boucle d'entraînement et garde le meilleur modèle sur le disque.

```bash
pip install yaxlib
```

## En trente secondes

Une sinusoïde bruitée, un perceptron multicouche, un entraînement :

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

X = jr.uniform(jr.key(0), (256, 1), minval=-3.0, maxval=3.0)
Y = jnp.sin(X) + 0.1 * jr.normal(jr.key(1), X.shape)

model = yax.layers.MLP(layer_sizes=(1, 32, 32, 1), activation="tanh", rkey=jr.key(2))
config = yax.configs.TrainConfig(learning_rate=1e-2, n_epoch=100)

run = yax.training.train(model, "out/sinus", config,
                         (X[:200], Y[:200], 32),     # entraînement, par batchs de 32
                         (X[200:], Y[200:]),         # validation
                         rkey=jr.key(3))

print(run.loss)                                      # perte de validation du meilleur modèle
Y_pred = yax.batch_apply(run.trained_model, X, None)
```

Le dossier `out/sinus/0` contient le meilleur modèle, sa configuration et
l'historique des pertes : `yax.training.load_run("out/sinus/0")` les relit
dans une autre session.

## Ce que yax vous donne

- **Un modèle est un pytree.** Déclarez une classe qui hérite de `yax.Module`,
  ses champs sont les paramètres, sa méthode `apply(x, rkey=None)` calcule la
  sortie pour un exemple. Rien à envelopper, rien à filtrer : `jax.grad(loss)(model)`
  rend un modèle de même forme dont les feuilles sont les gradients, et
  `optimizer.init(model)` l'accepte tel quel.
- **Une seule signature, partout.** `apply` traite un exemple, `jax.vmap` fait le
  batch ; `rkey` est la source d'aléatoire quand il y en a (dropout, tirages),
  et le mode d'évaluation se bascule d'un mot, `model.set_inference(True)`.
- **Un entraînement qui laisse des traces.** Chaque appel à `yax.training.train`
  crée un dossier numéroté avec le meilleur modèle, l'état de l'optimiseur, la
  configuration et l'historique. Les runs d'un même dossier se comparent, et
  `yax.training.find_best_run` désigne le meilleur.
- **Des briques prêtes à l'emploi**, toutes écrites comme vous les écririez :
  couches (MLP, convolutions n-dimensionnelles, GRU et LSTM, attention
  multi-têtes, blocs transformer, message passing), modèles complets (U-Net,
  opérateur de Fourier, VAE, flot normalisant, diffusion, mini-YOLO), pertes,
  optimiseurs d'optax par leur nom, prétraitements de données.
- **Des mini-fonctionnalités** sympathiques. Essayez par exemple `yax.ipprint(model)` dans un notebook.

## Pourquoi pas flax ou equinox ?

Par simplicité. Chaque framework se mesure aux concepts qu'il ajoute à jax :
flax apporte 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. yax ajoute le strict minimum pour construire des modèles à 
base de layers imbriqués. Tout ce qui n'est pas dans yax se code en jax
ordinaire, sans friction. 


## Pour continuer

- [Prise en main](https://octaviogame.com/recherche/yax/tutoriels/prise_en_main/) :
  un quart d'heure, des données, un modèle, un entraînement, un rechargement.
- [Aller plus loin](https://octaviogame.com/recherche/yax/tutoriels/aller_plus_loin/) :
  sous le capot du modèle, l'aléa et le mode inférence, les samplers, la
  configuration de l'entraînement, les pertes à soi.
- [Référence de l'API](https://octaviogame.com/recherche/yax/api/core/) et
  [démos](https://octaviogame.com/recherche/yax/exemples/demos/) : une
  démonstration par famille de modèles, qui converge en quelques secondes sur CPU.

## Développement

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