Checkpoints and Pretrained Models
A checkpoint is an immutable collection of tensors keyed by distinct, non-empty names, stored as a safetensors file. Typed parameter structures enter and leave checkpoints through names you declare; foreign checkpoints — HuggingFace Hub exports, say — are adapted checkpoint-to-checkpoint until they match your names. This guide covers both directions.
Named Structures
Checkpoint consumes Checkpoint.Named modules: Nx.Ptree.S plus one function, names, giving each tensor leaf a stable name in traversal order. By convention leaves are named after record fields, with nested structures joined by ".":
open Kaun
module Mlp = struct
type t = { l1 : Linear.t; l2 : Linear.t }
let map (f : 'a 'b. ('a, 'b) Nx.t -> ('a, 'b) Nx.t) { l1; l2 } =
{ l1 = Linear.map f l1; l2 = Linear.map f l2 }
let map2 (f : 'a 'b. ('a, 'b) Nx.t -> ('a, 'b) Nx.t -> ('a, 'b) Nx.t) p q =
{ l1 = Linear.map2 f p.l1 q.l1; l2 = Linear.map2 f p.l2 q.l2 }
let iter (f : 'a 'b. ('a, 'b) Nx.t -> unit) { l1; l2 } =
Linear.iter f l1;
Linear.iter f l2
let names { l1; l2 } =
List.map (( ^ ) "l1.") (Linear.names l1)
@ List.map (( ^ ) "l2.") (Linear.names l2)
let apply p x = Linear.apply p.l2 (Fn.relu (Linear.apply p.l1 x))
end
Each layer module ships its own names (Linear.names is ["w"; "b"], or ["w"] without a bias), so a model's names is the same one-liner shape as its traversals. For structures only known at runtime, Checkpoint.Ptree names dynamic Rune.Ptree.t trees by their path from the root.
Saving and Loading
of_params turns a structure into named entries; save writes them. Loading is template-based: construct the model first, then to_params ~like replaces its values with the file's entries of the same names — the template supplies structure, names, dtypes, and shapes, and its values are discarded:
let () =
Nx.Rng.run ~seed:0 @@ fun () ->
let init () =
{
Mlp.l1 = Linear.init ~inputs:4 ~outputs:8;
l2 = Linear.init ~inputs:8 ~outputs:2;
}
in
let params = init () in
let path = Filename.temp_file "kaun-doc" ".safetensors" in
Checkpoint.save path
(Checkpoint.of_params (module Mlp) ~prefix:"model" params);
let ckpt = Checkpoint.load path in
List.iter print_endline (Checkpoint.names ckpt);
(* model.l1.b, model.l1.w, model.l2.b, model.l2.w *)
let restored =
Checkpoint.to_params (module Mlp) ~prefix:"model" ~like:(init ()) ckpt
in
Sys.remove path;
(* The restored parameters equal the saved ones. *)
let x = Nx.randn Nx.float32 [| 2; 4 |] in
let d = Nx.max (Nx.abs (Nx.sub (Mlp.apply params x) (Mlp.apply restored x))) in
Printf.printf "max difference: %g\n" (Nx.item [] d)
A missing entry, shape mismatch, or dtype mismatch raises (~cast:true casts mismatched dtypes instead). Entries the template does not name are ignored — the basis for both multi-section files and partial loading.
One File, Several Sections
Because extraction ignores unnamed entries, one file holds model parameters, parameter-shaped optimizer state, and counters side by side, each under its own prefix. Saving and restoring full training state:
let () =
Nx.Rng.run ~seed:0 @@ fun () ->
let init () =
{
Mlp.l1 = Linear.init ~inputs:4 ~outputs:8;
l2 = Linear.init ~inputs:8 ~outputs:2;
}
in
let params = init () in
let ostate = Vega.adam_init (module Mlp) params in
let path = Filename.temp_file "kaun-doc" ".safetensors" in
Checkpoint.save path
(Checkpoint.concat
[
Checkpoint.of_params (module Mlp) ~prefix:"model" params;
Checkpoint.of_params (module Mlp) ~prefix:"optim.mu" ostate.mu;
Checkpoint.of_params (module Mlp) ~prefix:"optim.nu" ostate.nu;
Checkpoint.of_int "optim.step" ostate.step;
]);
(* Resuming: extract each section with its own prefix. *)
let ckpt = Checkpoint.load path in
let like = init () in
let params =
Checkpoint.to_params (module Mlp) ~prefix:"model" ~like ckpt
in
let ostate =
{
Vega.mu = Checkpoint.to_params (module Mlp) ~prefix:"optim.mu" ~like ckpt;
nu = Checkpoint.to_params (module Mlp) ~prefix:"optim.nu" ~like ckpt;
step = Checkpoint.to_int "optim.step" ckpt;
}
in
Sys.remove path;
ignore params;
Printf.printf "resumed at step %d\n" ostate.step
The optimizer moments checkpoint with the model's module because they have the model's shape — one more payoff of parameter-shaped state. Batch_norm running statistics work the same way, under their own prefix with (module Model.Stats).
To load a file into a partially different model — a new head on a pretrained backbone, say — extract each sub-structure with its own module and prefix; entries for the parts you replace are simply never asked for.
Foreign Checkpoints: kaun.hf
The kaun.hf library fetches files from HuggingFace Hub repositories into a local cache and loads safetensors checkpoints — single-file or sharded — as Checkpoint.t values. Downloading shells out to curl (it must be on PATH); fetched files are cached, so only the first access touches the network.
let ckpt = Kaun_hf.load_checkpoint "gpt2" in
List.iter print_endline (Kaun.Checkpoint.names ckpt)
(* h.0.attn.c_attn.bias, h.0.attn.c_attn.weight, ..., wte.weight *)
Hub checkpoints name and lay out tensors by the exporting framework's conventions, which will not match your records. Three checkpoint-to-checkpoint combinators adapt them, composed with (|>):
rename f t— replace every entry namenbyf n; return names you do not care about unchanged.transpose name t— swap the last two axes of one entry. Use it on weights stored with the opposite orientation, such astorch.nn.Linear'soutputs × inputsweights when your model expectsinputs × outputs.split name ~into t— replace one entry by equal sections along an axis. Use it on fused projections.
Leftover entries after adaptation are harmless: to_params ignores entries its template does not name.
The GPT-2 Story
examples/04-gpt2 runs the whole pipeline: it defines GPT-2 as a record of kaun layers (~150 lines), loads the real weights, and generates text. The adaptation is instructive because the HF checkpoint differs from the natural kaun model in two ways:
- Fused attention projections. HF stores each block's query, key, and value projections as one
c_attntensor of shape[n_embd; 3 * n_embd]; the model has three separateLinearlayers.splitcuts the fused weight and bias into thirds:
let split_qkv ckpt i =
let fused leaf = Printf.sprintf "h.%d.attn.c_attn.%s" i leaf in
let ours p leaf = Printf.sprintf "blocks.%d.attn.%s.%s" i p leaf in
ckpt
|> Hf.split (fused "weight")
~into:[ ours "q" "w"; ours "k" "w"; ours "v" "w" ]
|> Hf.split (fused "bias")
~into:[ ours "q" "b"; ours "k" "b"; ours "v" "b" ]
- Foreign names. Everything else is a pure renaming —
wte.weighttowte.table,h.0.ln_1.weighttoblocks.0.ln1.gamma, and so on — one total function passed torename. (GPT-2'sConv1Dweights are alreadyinputs × outputs, so no transposes are needed; a PyTorchnn.Linearexport would need them.) Entries the model does not use, like attention mask buffers, are left in place and ignored by extraction.
Adapted, the checkpoint matches the model's own names, and typed parameters come out through the ordinary template-based extraction:
let params =
Kaun_hf.load_checkpoint "gpt2"
|> Gpt2.of_hf ~n_layer:cfg.n_layer
|> Checkpoint.to_params (module Gpt2.Params) ~like:(Gpt2.make cfg) ~cast:true
The ~like template is a zero-initialized model built from the downloaded config.json (Kaun_hf.load_config fetches and parses it). There is no per-architecture loader in the library: the adaptation is ~40 lines of user code, and the same three combinators cover other exports.
Next Steps
- PyTorch Comparison —
state_dict,torch.save, andfrom_pretrainedin kaun terms - Layers and Models — where
namescomes from