Skip to main content

YA+ — le logo de yaxlib : un yack, un A et un +, imbriques

yaxlib

Un mini-framework de réseaux de neurones pour jax. Documentation — distribution yaxlib, import yax :

pip install yaxlib          # ou : pip install <url du wheel>
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

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

Download files

Download the file for your platform. If you're not sure which to choose, learn more about installing packages.

Source Distribution

yaxlib-0.1.6.tar.gz (150.9 kB view details)

Uploaded Source

Built Distribution

If you're not sure about the file name format, learn more about wheel file names.

yaxlib-0.1.6-py3-none-any.whl (59.4 kB view details)

Uploaded Python 3

File details

Details for the file yaxlib-0.1.6.tar.gz.

File metadata

  • Download URL: yaxlib-0.1.6.tar.gz
  • Upload date:
  • Size: 150.9 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? No
  • Uploaded via: uv/0.11.15 {"installer":{"name":"uv","version":"0.11.15","subcommand":["publish"]},"python":null,"implementation":{"name":null,"version":null},"distro":{"name":"macOS","version":null,"id":null,"libc":null},"system":{"name":null,"release":null},"cpu":null,"openssl_version":null,"setuptools_version":null,"rustc_version":null,"ci":null}

File hashes

Hashes for yaxlib-0.1.6.tar.gz
Algorithm Hash digest
SHA256 767a1ec51f8a06dc49c8212779a51e1189274096d25c37be4fdc232580e1a676
MD5 0f090041676d12124360a0366e9fa80f
BLAKE2b-256 81c6447a473a92887557d8f24b7921671d862057b74b3e8fbe66599330c219e5

See more details on using hashes here.

File details

Details for the file yaxlib-0.1.6-py3-none-any.whl.

File metadata

  • Download URL: yaxlib-0.1.6-py3-none-any.whl
  • Upload date:
  • Size: 59.4 kB
  • Tags: Python 3
  • Uploaded using Trusted Publishing? No
  • Uploaded via: uv/0.11.15 {"installer":{"name":"uv","version":"0.11.15","subcommand":["publish"]},"python":null,"implementation":{"name":null,"version":null},"distro":{"name":"macOS","version":null,"id":null,"libc":null},"system":{"name":null,"release":null},"cpu":null,"openssl_version":null,"setuptools_version":null,"rustc_version":null,"ci":null}

File hashes

Hashes for yaxlib-0.1.6-py3-none-any.whl
Algorithm Hash digest
SHA256 6473f72828ddaecea2bf245889bfb27af9c8d11100c4fa50ddc900d9a58ab3e6
MD5 a30c70508ed5504253b1d0131f9bed53
BLAKE2b-256 53dcc22c8a0d175a7dd5a68d36eebe5b97eeac55cb630920d43c455c461cf54b

See more details on using hashes here.

Release history Release notifications | RSS feed

This release

0.1.6 This release

2 files

0.1.5

2 files

0.1.4

2 files

0.1.3

2 files

0.1.2

2 files

0.1.1

2 files

0.1.0

2 files

Anthropic, PBC Visionary sponsor Bloomberg Visionary sponsor Hudson River Trading Visionary sponsor Meta Visionary sponsor NVIDIA Visionary sponsor Microsoft Sustainability sponsor Depot Continuous Integration AWS Cloud computing and Security Sponsor Datadog Monitoring Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page