Aller au contenu

yax — le cœur

La classe de base des modèles et les fonctions qui opèrent sur un modèle. Ces noms s'écrivent directement sous yax.

Module

Classe de base de tous les modèles et de toutes les couches.

Un Module est un pytree jax. On déclare ses champs par annotation de classe, on les remplit dans __init__, et on écrit la méthode apply(x, rkey=None), qui calcule la sortie pour un exemple.

Les champs sont de deux sortes :

  • dynamiques (annotation seule) : les paramètres. Ils contiennent des tableaux jax, des sous-modules, ou des listes, tuples et dictionnaires de ceux-ci. jax.grad, jax.jit et les optimiseurs d'optax les voient ;
  • statiques (= yax.StaticField()) : tout le reste, tailles, options, fonctions d'activation. Ils décrivent la forme du modèle et ne sont pas entraînés.

Une valeur non autorisée dans un champ dynamique (un flottant, une chaîne…) déclenche une erreur dès la construction. Un module est immuable : pour changer un champ, utiliser yax.tree_at.

class Classifieur(yax.Module):
    mlp: yax.layers.MLP
    dim_in: int = yax.StaticField()

    def __init__(self, dim_in, rkey):
        self.mlp = yax.layers.MLP((dim_in, 32, 1), "relu", rkey)
        self.dim_in = dim_in

    def apply(self, x, rkey=None):
        return self.mlp.apply(x)

Attributs :

Nom Type Description
inference bool

mode inférence, False par défaut. Les couches aléatoires, comme le dropout, y sont désactivées. Se change avec set_inference.

Code source dans yax/core.py
128
class Module:
166
167
168
    inference: bool = StaticField(default_value=False)

    _initializing = False   # défaut de classe : les instances unflatten le gardent

set_inference

set_inference(value)

Renvoie une copie du modèle en mode inférence (True) ou entraînement (False).

Le changement s'applique à tous les sous-modules ; le modèle d'origine n'est pas modifié. yax.training.train gère ce mode automatiquement.

Paramètres :

Nom Type Description Défaut
value bool

True pour le mode inférence, False pour l'entraînement.

obligatoire
Code source dans yax/core.py
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
def set_inference(self: _M, value: bool) -> _M:
    """Renvoie une copie du modèle en mode inférence (`True`) ou entraînement (`False`).

    Le changement s'applique à tous les sous-modules ; le modèle d'origine
    n'est pas modifié. `yax.training.train` gère ce mode automatiquement.

    Args:
        value: `True` pour le mode inférence, `False` pour l'entraînement.
    """

    def reconstruit(x):
        if isinstance(x, Module):
            klass = type(x)
            obj = object.__new__(klass)
            for name in klass._yax_dynamics:
                object.__setattr__(obj, name, reconstruit(x.__dict__.get(name)))
            for name in klass._yax_statics:
                object.__setattr__(obj, name, x.__dict__.get(name))
            object.__setattr__(obj, "inference", bool(value))
            return obj
        if isinstance(x, list):
            return [reconstruit(v) for v in x]
        if isinstance(x, tuple):
            return tuple(reconstruit(v) for v in x)
        if isinstance(x, dict):
            return {k: reconstruit(v) for k, v in x.items()}
        return x

    return reconstruit(self)

StaticField

yax.StaticField(*, default_value=None)

Déclare un champ statique d'un yax.Module.

Un champ statique contient ce qui n'est pas un paramètre : une taille, une option, une fonction. Il fait partie de la structure du modèle et n'est pas entraîné. Il peut aussi contenir un tableau constant : un masque, un encodage de position, le calendrier de bruit d'une diffusion.

class Couche(yax.Module):
    taille: int = yax.StaticField()
    mode: str = yax.StaticField(default_value="somme")

Tableaux statiques : quelques Mo au plus. Un tableau statique n'est pas une entrée du programme compilé par jax.jit, mais une constante écrite dedans. Sa taille alourdit donc la compilation : rien de visible en dessous du Mo, une compilation plusieurs fois plus longue au-delà de la dizaine de Mo. Pour un gros tableau constant, des poids pré-entraînés par exemple, mieux vaut un champ dynamique, figé dans apply par jax.lax.stop_gradient (attention : la décroissance des poids d'adamw le modifierait quand même). Une fois compilé, en revanche, un tableau statique ne coûte rien de plus à chaque pas, au contraire : on ne calcule ni son gradient ni sa mise à jour.

Paramètres :

Nom Type Description Défaut
default_value Any

valeur prise par le champ si __init__ ne le remplit pas.

None
Code source dans yax/core.py
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
def StaticField(*, default_value: Any = None) -> Any:
    """Déclare un champ statique d'un `yax.Module`.

    Un champ statique contient ce qui n'est pas un paramètre : une taille, une
    option, une fonction. Il fait partie de la structure du modèle et n'est pas
    entraîné. Il peut aussi contenir un tableau constant : un masque, un
    encodage de position, le calendrier de bruit d'une diffusion.

    ```python
    class Couche(yax.Module):
        taille: int = yax.StaticField()
        mode: str = yax.StaticField(default_value="somme")
    ```

    **Tableaux statiques : quelques Mo au plus.** Un tableau statique n'est pas
    une entrée du programme compilé par `jax.jit`, mais une constante écrite
    dedans. Sa taille alourdit donc la compilation : rien de visible en
    dessous du Mo, une compilation plusieurs fois plus longue au-delà de la
    dizaine de Mo. Pour un gros tableau constant, des poids pré-entraînés par
    exemple, mieux vaut un champ dynamique, figé dans `apply` par
    `jax.lax.stop_gradient` (attention : la décroissance des poids d'`adamw`
    le modifierait quand même). Une fois compilé, en revanche, un tableau
    statique ne coûte rien de plus à chaque pas, au contraire : on ne calcule
    ni son gradient ni sa mise à jour.

    Args:
        default_value: valeur prise par le champ si `__init__` ne le remplit pas.
    """
    # une fonction qui renvoie Any, et non une classe : un vérificateur de
    # types accepte ainsi `taille: int = yax.StaticField()` (même procédé que
    # dataclasses.field et equinox.field)
    return _ChampStatique(default_value)

tree_at

yax.tree_at(where, pytree, replace)

Renvoie une copie d'un modèle dont certaines feuilles sont remplacées.

model2 = yax.tree_at(lambda m: m.layers[-1].bias, model, jnp.ones(3))

Paramètres :

Nom Type Description Défaut
where Callable

fonction qui, appliquée au modèle, renvoie la feuille à remplacer, ou un tuple de feuilles.

obligatoire
pytree T

le modèle d'origine, qui n'est pas modifié.

obligatoire
replace

la nouvelle valeur, ou un tuple de valeurs.

obligatoire

Renvoie :

Type Description
T

Le modèle modifié.

Lève :

Type Description
ValueError

si where ne désigne pas une feuille du modèle.

Code source dans yax/core.py
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
def tree_at(where: Callable, pytree: T, replace) -> T:
    """Renvoie une copie d'un modèle dont certaines feuilles sont remplacées.

    ```python
    model2 = yax.tree_at(lambda m: m.layers[-1].bias, model, jnp.ones(3))
    ```

    Args:
        where: fonction qui, appliquée au modèle, renvoie la feuille à
            remplacer, ou un tuple de feuilles.
        pytree: le modèle d'origine, qui n'est pas modifié.
        replace: la nouvelle valeur, ou un tuple de valeurs.

    Returns:
        Le modèle modifié.

    Raises:
        ValueError: si `where` ne désigne pas une feuille du modèle.
    """
    cibles = where(pytree)
    if not isinstance(cibles, (tuple, list)):
        cibles, replace = (cibles,), (replace,)
    if len(cibles) != len(replace):
        raise ValueError(f"tree_at : {len(cibles)} cible(s) mais "
                         f"{len(replace)} remplacement(s).")

    leaves, treedef = jtu.tree_flatten(pytree)
    nouvelles = list(leaves)
    for cible, valeur in zip(cibles, replace):
        indices = [i for i, l in enumerate(leaves) if l is cible]
        if len(indices) != 1:
            raise ValueError(
                "tree_at : cible introuvable ou ambigue — `where` doit rendre "
                "une feuille (un tableau) extraite du pytree lui-même.")
        nouvelles[indices[0]] = valeur
    return jtu.tree_unflatten(treedef, nouvelles)

batch_apply

yax.batch_apply(model, x, rkey=None)

Applique un modèle à un lot d'exemples.

Un modèle yax traite un exemple à la fois ; batch_apply le vectorise avec jax.vmap sur la première dimension de x.

y = yax.batch_apply(model, x)          # évaluation, sans aléa
y = yax.batch_apply(model, x, rkey)    # avec aléa (dropout...)

Paramètres :

Nom Type Description Défaut
model Module

le modèle.

obligatoire
x ArrayLike

le lot, de forme (nb, ...).

obligatoire
rkey Array | None

clé aléatoire, répartie entre les exemples ; None (défaut) pour une évaluation sans aléa, comme dans Module.apply.

None

Renvoie :

Type Description
Array

Les sorties du modèle pour chaque exemple, de forme (nb, ...).

Code source dans yax/core.py
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
def batch_apply(model: Module, x: ArrayLike, rkey: jax.Array | None = None) -> jax.Array:
    """Applique un modèle à un lot d'exemples.

    Un modèle yax traite un exemple à la fois ; `batch_apply` le vectorise
    avec `jax.vmap` sur la première dimension de `x`.

    ```python
    y = yax.batch_apply(model, x)          # évaluation, sans aléa
    y = yax.batch_apply(model, x, rkey)    # avec aléa (dropout...)
    ```

    Args:
        model: le modèle.
        x: le lot, de forme `(nb, ...)`.
        rkey: clé aléatoire, répartie entre les exemples ; `None` (défaut)
            pour une évaluation sans aléa, comme dans `Module.apply`.

    Returns:
        Les sorties du modèle pour chaque exemple, de forme `(nb, ...)`.
    """
    if rkey is None:
        return jax.vmap(model.apply, in_axes=(0, None))(x, None)  # type: ignore[attr-defined]
    return jax.vmap(model.apply)(x, jax.random.split(rkey, jax.numpy.shape(x)[0]))  # type: ignore[attr-defined]

pprint

yax.pprint(x)

Affiche l'arborescence d'un modèle, en texte.

Les tableaux sont résumés par leur type et leur forme (f32[2,16]) ; les champs statiques et les fonctions sont affichés en clair. Voir aussi ipprint, sa version dépliable pour les notebooks.

Paramètres :

Nom Type Description Défaut
x Any

un modèle ou un pytree quelconque.

obligatoire
Code source dans yax/core.py
387
388
389
390
391
392
393
394
395
396
397
def pprint(x: Any) -> None:
    """Affiche l'arborescence d'un modèle, en texte.

    Les tableaux sont résumés par leur type et leur forme (`f32[2,16]`) ; les
    champs statiques et les fonctions sont affichés en clair. Voir aussi
    `ipprint`, sa version dépliable pour les notebooks.

    Args:
        x: un modèle ou un pytree quelconque.
    """
    print(_pformat(x, 0))

ipprint

yax.ipprint(x)

Affiche l'arborescence d'un modèle sous forme dépliable, dans un notebook.

Chaque sous-module se déplie d'un clic et indique son nombre de paramètres ; un clic sur un tableau affiche ses valeurs. Hors notebook, ipprint se comporte comme pprint. Un modèle placé en dernière ligne d'une cellule s'affiche de la même façon.

Paramètres :

Nom Type Description Défaut
x Any

un modèle ou un pytree quelconque.

obligatoire
Code source dans yax/core.py
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
def ipprint(x: Any) -> None:
    """Affiche l'arborescence d'un modèle sous forme dépliable, dans un notebook.

    Chaque sous-module se déplie d'un clic et indique son nombre de
    paramètres ; un clic sur un tableau affiche ses valeurs. Hors notebook,
    `ipprint` se comporte comme `pprint`. Un modèle placé en dernière ligne
    d'une cellule s'affiche de la même façon.

    Args:
        x: un modèle ou un pytree quelconque.
    """
    try:
        from IPython.display import display, HTML
        get_ipython  # type: ignore[name-defined]  # noqa: F821 — n'existe que sous IPython
    except (ImportError, NameError):
        pprint(x)
        return
    display(HTML(_html_document(x)))