news_article.exe
📰
#Google

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

2026年10月2日1 次浏览来源:MarkTechPost 阅读原文

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...

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

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.

Copy CodeCopiedUse a different Browser

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 arrays dtype, which is a code path Kauldron runs on every batch. Without the two-line replacement below, which uses jaxs 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.

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.

Configuration systems usually go wrong when one value is needed in several places, and Kauldrons answer is cfg.ref. We point a warmup-cosine schedules 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.

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 models 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.

Kauldrons 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.

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.

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.

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.

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, s

> 分享: