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
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 Browserimport 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.printexc(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 arrayspec as arrayspec # jax >= 0.10.1 moved jax.src.prng, but etils <= 1.14.0 still reaches for it whenever it # inspects an array's dtype.
Kauldron calls that code on every batch, so without this two-line # patch a Trainer raises AttributeError before it finishes a single step.The replacement uses # jax's own public dtype API and is a no-op on older jax.if not hasattr(jax.src, "prng"): arrayspec.isjaxrandomdtype = lambda dt: jax.
dtypes.issubdtype(dt, jax.dtypes.prngkey) import numpy as np import optax import flax from flax import linen as nn import kauldron from kauldron import kd, konfig, kontext from kauldron.typing import Float, typechecked print(f" kauldron {kauldron.version} | jax {jax.version} | flax {flax.
version}" f" | optax {optax.version}") print(f" devices: {jax.devices()}") print("\n Kauldron's pitch is modularity: it is the glue, not the framework.
Four pieces do the work:") print(" konfig -> 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 configisadict(): with konfig.
imports(): import optax as coptax # looks like optax, builds ConfigDict instead cfg = coptax.adam(learningrate=0.003) print(f" cfg = {cfg}") print(f" type = {type(cfg).name}") print(f" qualname = {cfg.qualname!r} <- the call, stored as data") cfg.
learningrate = 1e-4 # configs are mutable optimizer = konfig.resolve(cfg) # ...until you resolve them print(f" after cfg.learningrate = 1e-4 -> resolve() gives {type(optimizer).name}") print("\n An arbitrarily complex optimizer is still just nested dicts:") chain = coptax.chain( coptax.
clipbyglobalnorm(1.0), coptax.scalebyadam(b2=0.99), coptax.scalebylearningrate(0.003), ) asjson = json.dumps(json.loads(chain.tojson()), indent=2) print(textwrap.indent(asjson, " ")) rebuilt = konfig.resolve(konfig.ConfigDict(json.loads(chain.
tojson()))) 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}" configisadict() 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(learningrate=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 configreferences(): with konfig.imports(): import optax as coptax from kauldron import kd as ckd cfg = ckd.train.Trainer() cfg.numtrainsteps = 1000 cfg.schedules = { "lr": coptax.
warmupcosinedecayschedule( initvalue=0.0, peakvalue=1e-3, warmupsteps=100, decaysteps=cfg.ref.numtrainsteps, # <- a reference, not the value 1000 ) } at1000 = konfig.resolve(cfg.schedules["lr"]) cfg.numtrainsteps = 200 # one edit...at200 = konfig.resolve(cfg.schedules["lr"]) # ...
and the schedule already knows print(f" {'progress':>10s} {'lr @ 1000 steps':>18s} {'lr @ 200 steps':>16s}") for frac in (0.1, 0.5, 0.9): print(f" {frac:>9.0%} {float(at1000(int(1000frac))):>18.6f}" f" {float(at200(int(200*frac))):>16.6f}") print("\n Without .
ref the schedule would have frozen 1000 into itself, and a sweep over") print(" numtrainsteps would have silently trained on the wrong decay curve.") return (f"lr at 90% of training: {float(at1000(900)):.6f} (1000 steps)" f" vs {float(at200(180)):.
6f} (200 steps)") configreferences() 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 decaysteps at cfg.ref.
numtrainsteps rather than 1000, then change numtrainsteps 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 datac
Related
相關文章
剛剛,iQOO掏出年度旗艦,自研電競芯片性能提升15%,首款電競平板也來了
作者 | 陳駿達 編輯 | 心緣 9月29日報道,剛剛,vivo旗下iQOO品牌發佈了年度旗艦手機iQOO 16,這臺手機搭載了第六代驍龍8超級至尊版,全球首發了由iQOO和三星顯示聯合研發的2K 165Hz三星珠峰屏,並基於iQOO搭建的“3+2遊戲技術版圖”,提升了手機在視效、操控、直播和跨端遊戲等維度的體驗。
抽“錦鯉”享美食!“點亮杭州 碰見好運”服務消費季活動啟動
本文作者: Nemo 2026-09-25 10:17 導語:據瞭解,圍繞“點亮杭州 碰見好運”主題,活動將在9月21日至10月7日期間發放百萬級消費券,覆蓋吃喝玩樂購。9月24日,“點亮杭州 碰見好運”服務消費季活動在西湖區天目裡國際街區正式啟動。
聚焦院外管理提質增效|《急性冠狀動脈綜合徵患者院外長期隨訪管理共識》更新研討,胸痛中心智慧全程管理行動項目正式啟動
近日,第八屆“儒道心學”心血管病學會議、第十屆滬魯心血管病專家論壇、第九屆日照心血管峰會在山東日照召開。由葛均波院士領銜,黃愷、蘇國海、李春潔等數十位心血管領域權威專家參與,會上完成兩大核心動作:一是召開《急性冠狀動脈綜合徵患者院外長期隨訪管理共識》更新研討會,專家集體錨定共識修訂的核心方向;二是胸痛中心智慧全程管理行動項目正式啟動,以專家共識為指引推進先行落地驗證。作為醫療 AI 賦能院外管理創新的先行者,訊飛醫療執行總裁鹿曉亮受邀參會,與學界、業界共同推動心血管院外管理向標準化、智能化、全週期階段邁進。
型別安全 AI Jev 編碼指南:型別決策、校準信心與系統一模型的推測性扇出
在本教學中,我們使用 TypeSafe AI 的第一個系統一模型 Jev,它完全不生成文字:我們向其發送一段程式狀態和一組型別問題,它會回傳選擇、分數和是/否機率,我們的程式碼可以直接根據這些結果進行分支。我們安裝官方 Python SDK,進行第一次呼叫,同時使用三種問題原語,並觀察狀態的形狀如何影響模型能得知的資訊。接著,我們從回傳的機率重新計算已發布的信心統計數據,衡量將十個問題合併為一次呼叫與分開十次呼叫相比的效益,並建立 API 設計所針對的模式:信心門控路由、權重保留在程式碼中的複合評分、型別函式呼叫,以及以模型能處理的方式進行計數。

別讓一部片子倒在交付前:SkyProduction 搶先首發短劇質檢
上傳成片和字幕,系統自動從字幕、音畫、內容底線三個維度逐集跑一遍,輸出可下載、可回填的質檢報告。一部短劇剪完、導出,團隊最不願面對的往往是下一步:逐集看片。字幕裡藏沒藏同音錯字?人物開口了,字幕跟得上嗎?中間有沒有閃一下的黑屏、突然炸響的背景音樂?

商湯發佈 SenseNova U1 Pro 正式版圖片創作模型,至高支持 8K 分辨率及特殊長寬比
作者:沁滄(實習) 責編:沁滄 評論: 9 月 21 日消息,商湯今日發佈 SenseNova U1 Pro 正式版模型,上線商湯小浣熊,和商湯日日新 API 服務。據介紹,SenseNova U1 Pro 模型可通過內生的圖文交錯思維鏈創作出內容準確、設計精美、生產可用的圖片素材,效果比肩海外頂尖模型。