Introduzione a Kauldron

Kauldron è una libreria di training basata su JAX, rilasciata da Google Research, che promette velocità di ricerca e modularità. L’obiettivo è consentire a ricercatori e sviluppatori di costruire esperimenti complessi senza dover modificare il codice di base, sfruttando quattro componenti chiave: konfig per le configurazioni pure, kontext per il wiring via stringhe, ktype (runtime shape checker) per la verifica delle forme e kd.train per il Trainer.

1. Installazione e patch di compatibilità

Il primo passo è installare la versione 1.4.2 di Kauldron tramite pip. Poiché JAX 0.10.1 ha spostato il modulo privato jax._src.prng, è necessario applicare una piccola patch che reindirizza le chiamate di etils al nuovo API pubblico. Senza questa correzione il Trainer solleva un AttributeError al primo batch.

import os

import sys

import json

import textwrap

import traceback

import subprocess

subprocess.run([sys.executable, "-m", "pip", "install", "-q", "kauldron==1.4.2"], check=True)

import jax

from etils.enp import arrayspec as array_spec

if not hasattr(jax._src, "prng"):

arrayspec.isjaxrandomdtype = lambda dt: jax.dtypes.issubdtype(dt, jax.dtypes.prng_key)

Dopo la patch, importiamo le librerie fondamentali: jax, numpy, optax, flax.linen e, naturalmente, kauldron con i suoi sotto‑moduli kd, konfig e kontext. Un breve print mostra le versioni correnti e i dispositivi disponibili (CPU in questo esempio).

2. Konfig: le configurazioni come alberi di dizionari

All’interno di un blocco konfig.imports(), l’importazione di una libreria (ad esempio optax) restituisce oggetti “fittizi” che, anziché creare istanze, costruiscono ConfigDict. Questi dizionari contengono il nome qualificato della chiamata e i parametri, consentendo di serializzare l’intera configurazione in JSON.

with konfig.imports():

import optax as coptax

cfg = coptax.adam(learning_rate=0.003)

print(cfg) # → ConfigDict con _qualname_='optax.adam'

cfg.learning_rate = 1e-4 # mutabilità

optimizer = konfig.resolve(cfg) # crea l’oggetto reale

La serializzazione è completa: anche una catena di ottimizzatori complessa si riduce a JSON e può essere ricostruita senza alcuna modifica nella libreria originale.

3. Riferimenti dinamici con cfg.ref

Un vantaggio cruciale di konfig è la possibilità di creare riferimenti a valori ancora non definiti. Nell’esempio seguente, il numero di passi di training è usato per costruire una schedule di decadimento del learning rate. Cambiando cfg.numtrainsteps dopo la creazione della schedule, quest’ultima si aggiorna automaticamente.

cfg = kd.train.Trainer()

cfg.numtrainsteps = 1000

cfg.schedules = {

"lr": coptax.warmupcosinedecay_schedule(

init_value=0.0,

peak_value=1e-3,

warmup_steps=100,

decaysteps=cfg.ref.numtrain_steps,

)

}

lr_1000 = konfig.resolve(cfg.schedules["lr"])

cfg.numtrainsteps = 200

lr_200 = konfig.resolve(cfg.schedules["lr"])

Questo meccanismo evita errori silenti in sweep di iper‑parametri, garantendo che ogni esperimento utilizzi la curva di decadimento corretta.

4. Kontekst: wiring via stringhe

Kontext collega i componenti del training (modello, loss, metriche, ottimizzatore) usando percorsi chiave sotto forma di stringa. Nessun modulo importa direttamente un altro, il che riduce drasticamente le dipendenze circolari. Per esempio, una loss può riferirsi al modello mediante il path "model/forward" senza importare il file del modello.

    • Separazione netta: il loss non conosce la classe del modello.
    • Ri‑usabilità: lo stesso loss può essere ri‑usato con modelli diversi.
    • Facilità di testing: è possibile mockare parti del grafo semplicemente cambiando la stringa di riferimento.

5. Runtime shape checker: tipi con assi nominati

Kauldron introduce un sistema di typing dinamico basato su Float['*b h w c']. Gli assi nominati (b batch, h altezza, w larghezza, c canali) vengono legati tra gli argomenti di funzioni e verificati al volo. Se la forma non corrisponde, il framework stampa un messaggio esplicito indicando l’assi che non è stato associato.

from kauldron.typing import Float, typechecked

@typechecked

def my_loss(pred: Float['b h w c'], target: Float['b h w c']):

return ((pred - target) ** 2).mean()

Questo controllo è particolarmente utile quando si sperimentano architetture con dimensioni dinamiche, evitando errori di broadcasting difficili