AI 服務暫時不可用,以下為來源正文,待恢復後補全翻譯。
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 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. Copy CodeCopiedUse a different Browser @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. Copy CodeCopiedUse a different Browser @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. Copy CodeCopiedUse a different Browser @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. Copy CodeCopiedUse a different Browser @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. Copy CodeCopiedUse a different Browser 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. [truncated for AI cost control]