• bitcoinBitcoin(BTC)$84,897.001.70%
  • ethereumEthereum(ETH)$2,715.110.98%
  • tetherTether(USDT)$1.000.03%
  • binancecoinBNB(BNB)$771.870.41%
  • rippleXRP(XRP)$1.500.54%
  • usd-coinUSDC(USDC)$1.000.01%
  • solanaSolana(SOL)$119.551.27%
  • tronTRON(TRX)$0.334431-0.99%
  • Figure HelocFigure Heloc(FIGR_HELOC)$1.02-0.71%
  • zcashZcash(ZEC)$1,335.02-5.81%
  • HyperliquidHyperliquid(HYPE)$88.30-1.31%
  • dogecoinDogecoin(DOGE)$0.093911-0.82%
  • chainlinkChainlink(LINK)$14.29-0.58%
  • moneroMonero(XMR)$547.160.21%
  • whitebitWhiteBIT Coin(WBT)$84.821.63%
  • USDSUSDS(USDS)$1.000.04%
  • cardanoCardano(ADA)$0.246447-0.36%
  • RainRain(RAIN)$0.012038-1.90%
  • leo-tokenLEO Token(LEO)$8.92-1.24%
  • stellarStellar(XLM)$0.219194-4.24%
  • nearNEAR Protocol(NEAR)$4.89-6.44%
  • bitcoin-cashBitcoin Cash(BCH)$307.300.45%
  • uniswapUniswap(UNI)$8.981.55%
  • litecoinLitecoin(LTC)$68.762.57%
  • Ethena USDeEthena USDe(USDE)$1.000.01%
  • suiSui(SUI)$1.180.85%
  • avalanche-2Avalanche(AVAX)$10.920.55%
  • CantonCanton(CC)$0.120645-4.37%
  • Blockchain USDBlockchain USD(USDB)$0.871,000.00%
  • daiDai(DAI)$1.00-0.01%
  • hedera-hashgraphHedera(HBAR)$0.102709-1.19%
  • USD1USD1(USD1)$1.000.03%
  • the-open-networkGram (prev. Toncoin)(GRAM)$1.574.84%
  • BitwayBitway(BTW)$1.416.22%
  • quant-networkQuant(QNT)$254.30-14.00%
  • BittensorBittensor(TAO)$304.521.27%
  • shiba-inuShiba Inu(SHIB)$0.0000060.66%
  • crypto-com-chainCronos(CRO)$0.0683361.37%
  • tether-goldTether Gold(XAUT)$4,153.16-0.10%
  • Global DollarGlobal Dollar(USDG)$1.000.01%
  • paypal-usdPayPal USD(PYUSD)$1.000.05%
  • aaveAave(AAVE)$176.219.30%
  • Pump.funPump.fun(PUMP)$0.005753-0.03%
  • okbOKB(OKB)$121.10-0.36%
  • Ripple USDRipple USD(RLUSD)$1.000.00%
  • EthenaEthena(ENA)$0.241288-8.63%
  • Circle USYCCircle USYC(USYC)$1.140.01%
  • OndoOndo(ONDO)$0.493614-2.20%
  • MemeCoreMemeCore(M)$1.050.03%
  • Ondo US Dollar YieldOndo US Dollar Yield(USDY)$1.150.01%
TradePoint.io
  • Main
  • AI & Technology
  • Stock Charts
  • Market & News
  • Business
  • Finance Tips
  • Trade Tube
  • Blog
  • Shop
No Result
View All Result
TradePoint.io
No Result
View All Result

A Coding Guide to Google Research’s Kauldron: Configs That Are Plain Data, Components Wired by String, and a JAX Trainer You Can Read End to End

October 2, 2026
in AI & Technology
Reading Time: 17 mins read
A A
A Coding Guide to Google Research’s Kauldron: Configs That Are Plain Data, Components Wired by String, and a JAX Trainer You Can Read End to End
ShareShareShareShareShare

In this tutorial, we implement Kauldron, the JAX training library from Google Research that describes itself as optimized for research velocity and modularity, and we take those two words literally by testing what they actually buy us. We install it, then spend the first half of the notebook on the three mechanisms that make Kauldron different from a stack of Flax and Optax: konfig, which turns an experiment into a tree of plain dictionaries that round-trip through JSON; kontext, which wires components together with string key paths so a loss never imports the model it scores; and the runtime shape checker, whose named axes bind across arguments and report what they were bound to when something does not match. We then write a custom loss and a custom metric in the shape the framework expects, train a real Trainer on synthetic in-memory data with no download and no accelerator, and monitor an inner layer of the model without editing the model. We finish by running a five-variant sweep in which every experiment differs by one config line, and by letting a training run checkpoint itself and resume where it stopped.

import os
import sys
import json
import textwrap
import traceback
import subprocess
 
RESULTS = {}
 
 
def banner(title):
    print("\n" + "=" * 78)
    print(title)
    print("=" * 78)
 
 
def section(name):
    def wrap(fn):
        def run(*a, **kw):
            banner(name)
            try:
                out = fn(*a, **kw)
                RESULTS[name] = out if isinstance(out, str) else "ok"
                return out
            except Exception as e:
                RESULTS[name] = f"SKIPPED / FAILED -> {type(e).__name__}: {e}"
                print(f"\n[!] {name} did not complete: {type(e).__name__}: {e}")
                traceback.print_exc(limit=3)
                return None
        return run
    return wrap
 
 
banner("0. Install Kauldron, and the one compatibility patch you need today")
subprocess.run([sys.executable, "-m", "pip", "install", "-q", "kauldron==1.4.2"], check=True)
 
import jax
from etils.enp import array_spec as _array_spec
 
# jax >= 0.10.1 moved `jax._src.prng`, but etils  your experiment IS a Python call tree, and that tree is a plain dict")
print("    kontext   -> parts are wired by string key paths, so they never import each other")
print("    ktyping   -> Float['*b h w c'] checked at runtime, with named axes bound across args")
print("    kd.train  -> Trainer: model + data + losses + metrics + optimizer, and nothing else")
print("\n  Everything below runs on a CPU runtime with no dataset download: the data is synthetic.")

We install Kauldron and apply the one compatibility patch the current release combination needs. jax 0.10.1 moved the private module jax._src.prng, and etils up to 1.14.0 still reaches for it whenever it inspects an array’s dtype, which is a code path Kauldron runs on every batch. Without the two-line replacement below, which uses jax’s own public dtype API and is a no-op on older versions, a Trainer raises AttributeError before it completes a single step. With it in place we import the four pieces that do the work: konfig for the config system, kontext for the wiring, the typing module for runtime shape checks, and kd.train for the Trainer itself. Everything afterwards runs on a CPU runtime, because the only dataset in this notebook is one we generate.

YOU MAY ALSO LIKE

Florida County Finds 11 Unpermitted Flock Cameras With Unidentified Owners

OpenAI Fires Three Employees Who Allegedly Shared Info With An External AI Safety Group

@section("1. A config is a call tree, and a call tree is a dict")
def config_is_a_dict():
    with konfig.imports():
        import optax as coptax                      # looks like optax, builds ConfigDict instead
 
    cfg = coptax.adam(learning_rate=0.003)
    print(f"  cfg          = {cfg}")
    print(f"  type         = {type(cfg).__name__}")
    print(f"  __qualname__ = {cfg.__qualname__!r}    resolve() gives {type(optimizer).__name__}")
 
    print("\n  An arbitrarily complex optimizer is still just nested dicts:")
    chain = coptax.chain(
        coptax.clip_by_global_norm(1.0),
        coptax.scale_by_adam(b2=0.99),
        coptax.scale_by_learning_rate(0.003),
    )
    as_json = json.dumps(json.loads(chain.to_json()), indent=2)
    print(textwrap.indent(as_json, "    "))
 
    rebuilt = konfig.resolve(konfig.ConfigDict(json.loads(chain.to_json())))
    print(f"  JSON -> ConfigDict -> resolve() -> {type(rebuilt).__name__}")
    print("  optax has no idea konfig exists. No base class, no registry, no decorator.")
    return f"optax.chain -> JSON -> {type(rebuilt).__name__}"
 
 
config_is_a_dict()

We start with konfig, because it is the piece the rest of the library is built on. Inside a konfig.imports() block, importing optax gives us something that looks and autocompletes like optax but builds configuration instead of objects, so optax.adam(learning_rate=0.003) returns a ConfigDict holding the qualified name of the call and its arguments rather than an optimizer. That config is mutable until konfig.resolve turns it into the real thing, and because it is only nested dictionaries, an arbitrarily complex optax.chain serialises to JSON and comes back as a working optimizer. The important part is what optax had to do to support this: nothing. There is no base class, no registry, and no decorator anywhere in optax, and the same applies to any library we configure this way.

@section("2. cfg.ref: change one number, everything downstream follows")
def config_references():
    with konfig.imports():
        import optax as coptax
        from kauldron import kd as ckd
 
    cfg = ckd.train.Trainer()
    cfg.num_train_steps = 1000
    cfg.schedules = {
        "lr": coptax.warmup_cosine_decay_schedule(
            init_value=0.0, peak_value=1e-3, warmup_steps=100,
            decay_steps=cfg.ref.num_train_steps,    # 10s}  {'lr @ 1000 steps':>18s}  {'lr @ 200 steps':>16s}")
    for frac in (0.1, 0.5, 0.9):
        print(f"  {frac:>9.0%}  {float(at_1000(int(1000*frac))):>18.6f}"
              f"  {float(at_200(int(200*frac))):>16.6f}")
    print("\n  Without .ref the schedule would have frozen 1000 into itself, and a sweep over")
    print("  num_train_steps would have silently trained on the wrong decay curve.")
    return (f"lr at 90% of training: {float(at_1000(900)):.6f} (1000 steps)"
            f" vs {float(at_200(180)):.6f} (200 steps)")
 
 
config_references()

Configuration systems usually go wrong when one value is needed in several places, and Kauldron’s answer is cfg.ref. We point a warmup-cosine schedule’s decay_steps at cfg.ref.num_train_steps rather than 1000, then change num_train_steps to 200 and resolve the schedule again. The learning rate curve reshapes itself, because the config stored a reference rather than a copy of the value. Without that indirection the schedule would have frozen 1000 into itself, and a sweep over the number of training steps would have quietly trained every variant on the wrong decay curve, which is the kind of bug that produces a plausible number and no error.

@section("3. kontext: parts are wired by string, so they never import each other")
def kontext_keys():
    import dataclasses
 
    ctx = {
        "batch": {"image": np.zeros((4, 8, 8, 3)), "label": np.arange(4)},
        "preds": {"logits": np.ones((4, 10)), "aux": [{"pos": np.zeros(3)}]},
    }
    print("  a context is just nested data; a key path reaches into it:")
    for path in ["batch.image", "preds.logits", "preds.aux[0].pos"]:
        print(f"    {path:22s} -> {kontext.get_by_path(ctx, path).shape}")
 
    try:
        kontext.get_by_path(ctx, "batch.nope")
    except KeyError as e:
        print(f"    {'batch.nope':22s} -> KeyError: {str(e)[:96]}...")
 
    @dataclasses.dataclass(eq=True, frozen=True, kw_only=True)
    class MeanGap:
        preds: kontext.Key = kontext.REQUIRED       # these ARE the wiring
        targets: kontext.Key = kontext.REQUIRED
 
        def __call__(self, *, preds, targets):
            return float(abs(np.asarray(preds).mean() - np.asarray(targets).mean()))
 
    metric = MeanGap(preds="preds.logits", targets="batch.label")
    kwargs = kontext.resolve_from_keyed_obj(ctx, metric)
    print(f"\n  MeanGap declared preds={metric.preds!r}, targets={metric.targets!r}")
    print(f"  resolved to kwargs: {{{', '.join(f'{k}: {v.shape}' for k, v in kwargs.items())}}}")
    print(f"  value = {metric(**kwargs)}")
    print("\n  MeanGap never imported the model and the model never heard of MeanGap. Point the")
    print("  same metric at 'preds.aux[0].pos' and nothing but that string changes.")
    return f"MeanGap(preds="preds.logits", targets="batch.label") = {metric(**kwargs)}"
 
 
kontext_keys()

kontext is how Kauldron connects components that know nothing about each other. A context is ordinary nested data, and a key path such as batch.image or preds.aux[0].pos reaches into it, resolving dictionary keys, attributes and list indices alike, and raising a KeyError that lists what was actually available when it cannot. Any object can declare its inputs by annotating fields as kontext.Key, and resolve_from_keyed_obj then pulls exactly those paths out of the context and hands them over as keyword arguments. We build a small metric this way and point it at a model’s outputs: the metric never imports the model, the model never hears of the metric, and redirecting the metric at a different tensor is a change to one string.

@section("4. ktyping: named axes, checked at runtime, bound across arguments")
def shape_checking():
    @typechecked
    def project(features: Float["*b n c"], weights: Float["c d"]) -> Float["*b n d"]:
        return jax.numpy.einsum("...c,cd->...d", features, weights)
 
    out = project(jax.numpy.zeros((2, 16, 8)), jax.numpy.zeros((8, 32)))
    print(f"  project(f32[2 16 8], f32[8 32]) -> {out.shape}   c bound to 8, d bound to 32")
 
    print("\n  now break it: c is bound to 8 by the first argument, so 5 cannot also be c")
    try:
        project(jax.numpy.zeros((2, 16, 8)), jax.numpy.zeros((5, 32)))
    except Exception as e:
        print(textwrap.indent(str(e), "    "))
    print("\n  'Inferred Dims' is the part worth having: it reports what each axis name was already")
    print("  bound to, so a mismatch names the axis instead of printing two anonymous shapes.")
    return "mismatch named the axis: c already bound to 8, got 5"
 
 
shape_checking()

Kauldron’s typing module checks array shapes at runtime using named axes. We annotate a function with Float[‘*b n c’] and Float[‘c d’], and the decorator binds each axis name the first time it sees it, then enforces that binding everywhere else in the signature, including the return value. When we deliberately pass an incompatible second argument, the error does the thing that matters: alongside the actual shapes it prints an Inferred Dims block showing that c had already been bound to 8, so the failure names the axis that disagreed instead of leaving us to compare two anonymous tuples. On a model with several tensors in flight this is the difference between a one-line fix and a debugging session.

import dataclasses
 
 
@dataclasses.dataclass(eq=True, frozen=True, kw_only=True)
class LogCosh(kd.losses.Loss):
    """log(cosh(err)): quadratic near zero, linear in the tails. ~30 lines less than raw Flax."""
 
    preds: kontext.Key = kontext.REQUIRED
    targets: kontext.Key = kontext.REQUIRED
 
    @typechecked
    def get_values(self, preds: Float["*a"], targets: Float["*a"]) -> Float["*a"]:
        return jax.numpy.log(jax.numpy.cosh(preds - targets))
 
 
@dataclasses.dataclass(eq=True, frozen=True, kw_only=True)
class WithinTol(kd.metrics.Metric):
    """Fraction of predictions landing within `tol` of the target, over every batch seen."""
 
    preds: kontext.Key = kontext.REQUIRED
    targets: kontext.Key = kontext.REQUIRED
    tol: float = 0.25
 
    @flax.struct.dataclass
    class State(kd.metrics.AutoState):
        # sum_field() marks a value that is ADDED when two states merge. Keeping the numerator
        # and the denominator apart is what makes the pooled result exact.
        n_hit: Float[""] = kd.metrics.sum_field(default=0.0)
        n_total: Float[""] = kd.metrics.sum_field(default=0.0)
 
        def compute(self) -> Float[""]:
            # Return a jax scalar, like the built-in states do: the metric writer that
            # `trainer.train()` logs through does not accept a bare numpy scalar.
            total = jax.numpy.maximum(jax.numpy.asarray(self.n_total), 1.0)
            return jax.numpy.asarray(self.n_hit) / total
 
    @typechecked
    def get_state(self, preds: Float["*a"], targets: Float["*a"]) -> "WithinTol.State":
        hit = (jax.numpy.abs(preds - targets) 4.0f} / {float(st.n_total):>3.0f}  = {float(st.compute()):.4f}")
    naive = (float(big.compute()) + float(small.compute())) / 2
    print(f"    {'mean of the two rates':22s} {'':>4s}   {'':>3s}  = {naive:.4f}   

We write a custom loss and a custom metric in the exact shape Kauldron expects, which is a frozen dataclass with kontext.Key fields and one method. The loss implements get_values and returns a per-element array; the framework handles reduction and the weight argument, which we confirm by checking that weight=0.5 halves the result exactly. The metric is more interesting, because a Kauldron metric does not return a number but a State that merges. We build ours on AutoState with two sum_field entries, a numerator and a denominator, and then merge a six-row batch with a two-row batch. The pooled result is exact, while the average of the two per-batch rates is visibly wrong, which is precisely what would happen on the ragged last batch of an epoch, and since merge is associative the same mechanism aggregates a metric across devices without caring what order the results arrive in.

SEED_W = np.random.default_rng(7).normal(size=(16, 1)).astype("float32")
 
 
def make_loader(split: str):
    """A dataset is any callable returning an array tree. No download, no tf.data, no TFDS."""
 
    def load():
        rng = np.random.default_rng(0 if split == "train" else 1)
        n = 512 if split == "train" else 128
        x = rng.normal(size=(n, 16)).astype("float32")
        y = (x @ SEED_W + 0.1 * rng.normal(size=(n, 1))).astype("float32")
        return {"x": x, "y": y}
 
    return load
 
 
class MLP(nn.Module):
    inputs: kontext.Key = kontext.REQUIRED          # the model declares where its input comes from
    hidden: int = 32
 
    @nn.compact
    def __call__(self, inputs: Float["*b f"]) -> dict[str, Float["*b 1"]]:
        h = nn.relu(nn.Dense(self.hidden, name="enc")(inputs))
        return {"y": nn.Dense(1, name="out")(h)}
 
 
@section("6. A real Trainer on synthetic data, and a look inside the model")
def train_for_real():
    train_ds = kd.data.InMemoryPipeline(
        loader=make_loader("train"), batch_size=32, shuffle=True, num_epochs=None, seed=0)
    print(f"  element_spec: {kd.inspect.json_spec_like(train_ds.element_spec)}")
    print("\n  batch statistics straight from kd.inspect:")
    print(textwrap.indent(str(kd.inspect.get_batch_stats(next(iter(train_ds)))), "    "))
 
    trainer = kd.train.Trainer(
        seed=0,
        workdir="/tmp/kauldron_tutorial",
        train_ds=train_ds,
        model=MLP(inputs="batch.x"),
        num_train_steps=300,
        train_losses={"logcosh": LogCosh(preds="preds.y", targets="batch.y")},
        train_metrics={
            "within_tol": WithinTol(preds="preds.y", targets="batch.y"),
            "enc_norm": kd.metrics.Norm(tensor="interms.enc.__call__[0]"),   # an INNER layer
        },
        optimizer=optax.adam(learning_rate=1e-2),
    )
 
    it = iter(trainer.train_ds)
    state = trainer.trainstep.init(trainer.train_ds.element_spec)
    print(f"\n  params: {jax.tree.map(lambda x: tuple(x.shape), state.params)}")
    print(f"\n  {'step':>5s} {'logcosh':>10s} {'within .25':>11s} {'enc_norm':>10s}")
    for step in range(1, 301):
        # aux is opt-in: the step skips building it unless you ask, because it costs device time.
        state, aux = trainer.trainstep.step(state, next(it), return_losses=True, return_metrics=True)
        if step in (1, 25, 50, 100, 200, 300):
            losses = {k: float(v.compute()) for k, v in aux.loss_states.items()}
            metrics = {k: float(v.compute()) for k, v in aux.metric_states.items()}
            first_loss = losses["logcosh"] if step == 1 else first_loss
            final_loss = losses["logcosh"]
            print(f"  {step:>5d} {losses['logcosh']:>10.4f} {metrics['within_tol']:>11.4f}"
                  f" {metrics['enc_norm']:>10.4f}")
 
    print("\n  enc_norm was never returned by the model. 'interms.enc.__call__[0]' reaches into the")
    print("  Dense layer named 'enc' through Flax's captured intermediates, so monitoring an inner")
    print("  activation costs one string in the config and zero edits to MLP.")
    return f"logcosh {first_loss:.4f} -> {final_loss:.4f} over 300 CPU steps"
 
 
train_for_real()

Now we train. A Kauldron dataset is any callable returning a tree of arrays, so kd.data.InMemoryPipeline turns our synthetic regression data into a real pipeline with batching and shuffling and nothing to download. We assemble a Trainer from a model, that pipeline, our custom loss, our custom metric and an Optax optimizer, and drive its train step directly so we can print a loss curve: the loss falls from 1.51 to 0.005 over three hundred steps in about a second of CPU time. Two details are worth keeping. The train step does not build its auxiliary outputs unless we ask with return_losses and return_metrics, because computing them costs device time. And the enc_norm column is read from interms.enc.__call__[0], a key path into the intermediate output of the Dense layer named enc, so monitoring an inner activation costs one string in the config and no edit at all to the model.

@section("7. A sweep is a for-loop over config overrides")
def sweep():
    with konfig.imports():
        import optax as coptax
        from kauldron import kd as ckd
    with konfig.imports(lazy=True):
        # Classes defined in a notebook live in __main__, which cannot be fake-imported eagerly.
        from __main__ import MLP as CfgMLP
        from __main__ import make_loader as cfg_make_loader
        from __main__ import LogCosh as CfgLogCosh
 
    def base_config():
        cfg = ckd.train.Trainer()
        cfg.seed = 0
        cfg.workdir = "/tmp/kauldron_sweep"
        cfg.train_ds = ckd.data.InMemoryPipeline(
            loader=cfg_make_loader("train"), batch_size=32, shuffle=True, num_epochs=None)
        cfg.model = CfgMLP(inputs="batch.x", hidden=32)
        cfg.num_train_steps = 200
        cfg.train_losses = {"logcosh": CfgLogCosh(preds="preds.y", targets="batch.y")}
        cfg.optimizer = coptax.adam(learning_rate=1e-2)
        return cfg
 
    print(f"  cfg.model     = {base_config().model}")
    print(f"  cfg.optimizer = {base_config().optimizer}")
    print("\n  A bare lazy-imported name is a reference, not a call:")
    print(f"    CfgMLP(inputs=...)  -> {{'__qualname__': '__main__.MLP', ...}}   (built when resolved)")
    print(f"    cfg_make_loader     -> {dict(cfg_make_loader)}   (handed over as-is)")
 
    print(f"\n  {'override':34s} {'final logcosh':>14s}")
    results = {}
    for label, apply_override in [
        ("(baseline)", lambda c: None),
        ("cfg.model.hidden = 4", lambda c: setattr(c.model, "hidden", 4)),
        ("cfg.model.hidden = 128", lambda c: setattr(c.model, "hidden", 128)),
        ("cfg.optimizer.learning_rate = 0.1", lambda c: setattr(c.optimizer, "learning_rate", 0.1)),
        ("cfg.optimizer = optax.sgd(0.05)", lambda c: setattr(c, "optimizer", coptax.sgd(0.05))),
    ]:
        cfg = base_config()
        apply_override(cfg)
        trainer = konfig.resolve(cfg)               # ConfigDict -> a real, frozen Trainer
        it = iter(trainer.train_ds)
        state = trainer.trainstep.init(trainer.train_ds.element_spec)
        for _ in range(200):
            state, aux = trainer.trainstep.step(state, next(it), return_losses=True)
        results[label] = float(aux.loss_states["logcosh"].compute())
        print(f"  {label:34s} {results[label]:>14.4f}")
 
    print("\n  Five experiments, five one-line edits, and not one character of MLP, LogCosh or the")
    print("  training loop changed. On the command line the same overrides are")
    print("  --cfg.model.hidden=128, which is why a Kauldron sweep is a list of these strings.")
    return (f"best {min(results.values()):.4f} ({min(results, key=results.get)}),"
            f" worst {max(results.values()):.4f} ({max(results, key=results.get)})")
 
 
sweep()

This step is the argument for the whole design. We write the experiment once as a config, then run five variants that differ by exactly one line each: two model widths, two optimizer settings, and a wholesale swap of Adam for SGD. Every variant resolves into a fresh Trainer and trains for two hundred real steps, and not one character of the model, the loss or the training loop changes between them. Two konfig details make this work in a notebook. Classes defined in a notebook live in __main__, which cannot be fake-imported eagerly, so we use konfig.imports(lazy=True) for them. And a bare lazy-imported name resolves to the object itself rather than calling it, which is how the loader function is handed to the pipeline intact. On the command line these same overrides are written –cfg.model.hidden=128, which is why a Kauldron sweep is nothing more than a list of such strings.

@section("8. The guardrail: configs hold configs, never resolved objects")
def guardrails():
    with konfig.imports():
        from kauldron import kd as ckd
 
    cfg = ckd.train.Trainer()
    print("  Assigning a REAL flax module into a config is refused on the spot:")
    try:
        cfg.model = MLP(inputs="batch.x", hidden=8)
    except ValueError as e:
        print(textwrap.indent(str(e)[:420], "    "))
 
    print("\n  Why this matters: a half-resolved config cannot be serialized, diffed or overridden")
    print("  from the command line, so konfig refuses to let one exist rather than failing later.")
 
    print("\n  The same discipline shows up in sub-objects, which default to root-config references:")
    print(f"    kd.data.InMemoryPipeline(...).seed   default -> _FakeRootCfg('cfg.seed')")
    print(f"    kd.evals.Evaluator(...).ds           default -> _FakeRootCfg('cfg.eval_ds')")
    print("  Inside a Trainer those are filled from the root. Built standalone they are not, which")
    print("  is why step 6 passed seed=0 to the pipeline explicitly.")
    return "ConfigDict refused a resolved flax module at assignment"
 
 
guardrails()

Kauldron refuses to let a config hold a resolved object, and the refusal is immediate rather than deferred. Assigning a real Flax module into a ConfigDict raises on the spot with a message that suggests the two fixes, wrapping the import in konfig.imports() or using mock_modules in a notebook. A half-resolved config cannot be serialized, diffed, or overridden from the command line, so the system rules out the state entirely instead of failing later in a way that is hard to trace. The same discipline explains something we met in step 6: sub-objects such as pipelines and evaluators default their seed and dataset fields to references into the root config, which are filled in when the object is built inside a Trainer and are not when it is built standalone, which is why we passed the pipeline an explicit seed.

@section("9. Evaluation and checkpointing, and picking up where you left off")
def eval_and_checkpoint():
    import pathlib
    import shutil
 
    workdir = "/tmp/kauldron_resume"
    shutil.rmtree(workdir, ignore_errors=True)      # start from a clean slate for the demo
 
    def build(num_steps):
        return kd.train.Trainer(
            seed=0,
            workdir=workdir,
            train_ds=kd.data.InMemoryPipeline(
                loader=make_loader("train"), batch_size=32, shuffle=True, num_epochs=None, seed=0),
            model=MLP(inputs="batch.x"),
            num_train_steps=num_steps,
            log_metrics_every=100,
            train_losses={"logcosh": LogCosh(preds="preds.y", targets="batch.y")},
            train_metrics={"within_tol": WithinTol(preds="preds.y", targets="batch.y")},
            optimizer=optax.adam(learning_rate=1e-2),
            checkpointer=kd.ckpts.Checkpointer(save_interval_steps=100),
            evals={
                "eval": kd.evals.Evaluator(
                    run=kd.evals.EveryNSteps(100),
                    ds=kd.data.InMemoryPipeline(
                        loader=make_loader("eval"), batch_size=32, shuffle=False,
                        num_epochs=1, seed=0),
                    num_batches=4,
                )
            },
        )
 
    state, _ = build(200).train()
    print(f"  first run finished at step {int(state.step)}")
    saved = sorted(p.name for p in pathlib.Path(workdir).glob("checkpoints/ckpt_*"))
    print(f"  checkpoints on disk: {saved}")
 
    print("\n  Now build the same Trainer again, on the same workdir, asking for more steps:")
    state2, _ = build(300).train()
    print(f"  second run finished at step {int(state2.step)}")
    print("  The progress bar above started at 200, not 0: train() found the checkpoint and")
    print("  resumed, which is also what happens when a preemptible job is restarted.")
    print("\n  The evaluator ran on its own dataset every 100 steps. It inherited the model, the")
    print("  losses and the metrics from the root config, so declaring it took four lines.")
    return f"trained, evaluated, checkpointed and resumed at step {int(state2.step)}"
 
 
eval_and_checkpoint()

We close the loop with the parts that turn a training script into a job. An evaluator is declared in four lines, because it inherits the model, the losses and the metrics from the root config and only needs its own dataset and a schedule saying how often to run. A checkpointer writes state at a fixed step interval. Then we build the same Trainer a second time against the same working directory, asking for more steps, and the progress bar starts at 200 rather than 0: the run found its checkpoint and resumed. That is the same path taken when a preemptible job is restarted, and it is worth seeing once in a notebook where the whole thing takes a second, rather than discovering it for the first time on a cluster.

banner("SUMMARY")
for name, res in RESULTS.items():
    print(f"  {name:

The summary prints the one-line result each section returned, then points at where to go next: the two self-contained modules worth reading on their own, the real data pipelines that replace our synthetic one, the example configs in the repository, and the command line that turns the sweep from step 7 into flags.

In conclusion, we treated Kauldron’s two claims as things to test rather than repeat, and both held up for an easy-to-state reason. Modularity here is not an abstraction layer but the absence of one: a config is a dictionary describing a Python call, wiring is a string naming a path through data, and neither optax nor our own model needed a single line of Kauldron-aware code to participate. Research velocity follows from that, and we saw it concretely in the sweep, where five experiments cost five edited lines, and in the shape checker, which names the axis that disagreed rather than printing two shapes and leaving us to work it out. The parts we did not need, TensorFlow datasets, Grain, sharding and XManager, sat entirely out of the way while a Trainer ran on synthetic data on a CPU. What we would carry into real work is smaller than the library: keep configuration as data, wire components by path rather than by import, and make metrics states that merge, since all three are useful even in a codebase that never adopts Kauldron itself.


Check out the FULL CODES here. All credit goes to the researcher of this project. Also, feel free to follow us on Twitter and don’t forget to join our 150k+ML SubReddit and Subscribe to our Newsletter. Wait! are you on telegram? now you can join us on telegram as well.

Need to partner with us for promoting your GitHub Repo OR Hugging Face Page OR Product Release OR Webinar etc.? Connect with us


Sana Hassan, a consulting intern at Marktechpost and dual-degree student at IIT Madras, is passionate about applying technology and AI to address real-world challenges. With a keen interest in solving practical problems, he brings a fresh perspective to the intersection of AI and real-life solutions.

Credit: Source link

ShareTweetSendSharePin

Related Posts

Florida County Finds 11 Unpermitted Flock Cameras With Unidentified Owners
AI & Technology

Florida County Finds 11 Unpermitted Flock Cameras With Unidentified Owners

October 1, 2026
OpenAI Fires Three Employees Who Allegedly Shared Info With An External AI Safety Group
AI & Technology

OpenAI Fires Three Employees Who Allegedly Shared Info With An External AI Safety Group

October 1, 2026
Microsoft Launches MAI-Transcribe-2-Streaming and Two MAI-Voice Models – Unite.AI
AI & Technology

Microsoft Launches MAI-Transcribe-2-Streaming and Two MAI-Voice Models – Unite.AI

October 1, 2026
Judge Dismisses Lawsuits Claiming Google’s AI Overviews Siphon Web Traffic
AI & Technology

Judge Dismisses Lawsuits Claiming Google’s AI Overviews Siphon Web Traffic

October 1, 2026

Leave a Reply Cancel reply

Your email address will not be published. Required fields are marked *

Search

No Result
View All Result
Police visit home of Hayden Panettiere’s boyfriend

Police visit home of Hayden Panettiere’s boyfriend

September 27, 2026
Cricut’s New DIY Machines Let You Print And Cut Your Own Stickers

Cricut’s New DIY Machines Let You Print And Cut Your Own Stickers

September 25, 2026
Current with Christine Romans – Aug. 21 | NBC News NOW

Current with Christine Romans – Aug. 21 | NBC News NOW

September 26, 2026

About

Learn more

Our Services

Legal

Privacy Policy

Terms of Use

Bloggers

Learn more

Article Links

Contact

Advertise

Ask us anything

©2020- TradePoint.io - All rights reserved!

Tradepoint.io, being just a publishing and technology platform, is not a registered broker-dealer or investment adviser. So we do not provide investment advice. Rather, brokerage services are provided to clients of Tradepoint.io by independent SEC-registered broker-dealers and members of FINRA/SIPC. Every form of investing carries some risk and past performance is not a guarantee of future results. “Tradepoint.io“, “Instant Investing” and “My Trading Tools” are registered trademarks of Apperbuild, LLC.

This website is operated by Apperbuild, LLC. We have no link to any brokerage firm and we do not provide investment advice. Every information and resource we provide is solely for the education of our readers. © 2020 Apperbuild, LLC. All rights reserved.

No Result
View All Result
  • Main
  • AI & Technology
  • Stock Charts
  • Market & News
  • Business
  • Finance Tips
  • Trade Tube
  • Blog
  • Shop

© 2023 - TradePoint.io - All Rights Reserved!