Metadata-Version: 2.5
Name: yaxlib
Version: 0.3.4
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: mypy; 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 des réseaux de neurones avec jax, sans concepts ajoutés (bio). 
Un modèle est une classe Python dont les
tableaux sont les paramètres : `jax.grad`, `jax.jit`, `jax.vmap` et la fameuse lib d'optimisation optax
s'appliquent sur les modèles directement. Si on le souhaite, une fonction `yax.training.train`   fait la
boucle d'entraînement et sauvegarde le meilleur modèle.

```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, nb_epochs=100)

run = yax.training.train("out/sinus", config,
                         yax.obj.mse,                # l'objectif à minimiser
                         model,
                         (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)
```

Le dossier `out/sinus/0` contient l'historique de l'entrainement ainsi que le modèle entrainé, sauvegardé à sa meilleure validation.
Dans une autre session, `yax.training.load_run("out/sinus/0")` permet de retrouver tout cela.
Un second entrainement (en variant la config par exemple) s'enregistrera dans `"out/sinus/1"` etc.


## 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 entrainables ou statiques. Rien à envelopper, rien à filtrer : `jax.grad(loss)(model)`
  rend un modèle de même forme dont les feuilles sont les gradients.
- **Une seule signature, partout.** `apply(x,rkey)` traite un exemple `x`. Ensuite `jax.vmap` permet de traiter un batch;
  `rkey` est la source d'aléatoire quand il y en a (dropout, tirages).
   Le mode d'évaluation se bascule avec la méthode: `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 au sein du `mother_folder`.
- **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), fonctions de pertes,
  optimiseurs, 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/
```
