Skip to content

Safetensors for every dtype, atomic checkpoints, and the token stream

ModuleL0.6 · build · Python · Pass 2 · 5 to 7 h
You buildyou take over python/tinyllm/io/safetensors.py from L0.0 and add F16, BF16, F8_E4M3, I8, U8, I32 plus read_header; python/tinyllm/io/checkpoint.py: save_checkpoint, load_checkpoint, verify_step_dir; python/tinyllm/io/tokens.py: read_tokens_header, open_tokens, TokenStream
Contractcourse/contracts/py/tinyllm/io/safetensors.pyi · course/contracts/py/tinyllm/io/checkpoint.pyi · course/contracts/py/tinyllm/io/tokens.pyi · formats: safetensors.md, checkpoint.md, tokens-bin.md
Testscourse/tests/L0.6/ (what they check: section 4), golden files from the pinned safetensors library in course/fixtures/L0.6/safetensors/; L0.0’s safetensors tests keep running as your regression suite
NeedsM09.1 the BF16 bit converters · L0.1, L0.2, L0.4 the model the checkpoint tests train · M10.2 SGD and M10.3 AdamW (optimizer state to save) · M06.3 the PCG32 whose state the token cursor carries · L0.5 its bigram.py (L0.0’s tests, your regression suite, exercise it) · reading: L0.0 (or --ref-deps)
Used byL10.0 your engine reads the file this writer produces (inherited from L0.0 with safetensors.py) · L2.1 the n-gram model saves and loads its tables with save_safetensors and load_safetensors · later: L2.2 reads token streams, L7.9 loads BF16 HF weights, L10.1 memory-maps the same layout, the capstone trainer and dur.09 resume from these checkpoints · later: L3.6, L4.1, L5.5, L6.1, L6.2, L6.3, L6.5, L6.7
MilestoneMS-L0 (step 4: a run killed after step 60 resumes from step 50 and ends bitwise equal to an uninterrupted one)
Optional depththe safetensors README and safetensors/src/tensor.rs; Micikevicius et al., “FP8 Formats for Deep Learning” (2022); Pillai et al., “All File Systems Are Not Created Equal” (OSDI 2014) on crash consistency
  • One safetensors writer for seven dtypes, byte-identical to the reference library: tensors laid out by dtype first, then by name (test_matches_library_bytes, test_dtype_order_then_name).
  • BF16 and F8_E4M3 have no numpy type, so they travel as float32 values rounded to nearest, ties to even, or as raw bits written untouched; F8_E4M3 has no infinity, so a value beyond 448 is an error, not a saturation (test_bf16_rounds_to_nearest_even, test_f8_rounds_to_nearest_even, test_f8_rejects_out_of_range).
  • A checkpoint is written into a .tmp directory, every file fsynced, the manifest last, then renamed into place, then LATEST replaced by a rename: a crash at any instant leaves the old checkpoint or the new one (test_crash_at_every_write_keeps_a_valid_checkpoint).
  • Loading checks every file against the manifest and falls back to the newest older checkpoint that verifies (test_load_skips_incomplete_and_corrupt).
  • The token stream’s cursor (shard, offset, generator state) is its whole position: N batches, save, restore, M batches equals N + M batches (test_cursor_restore_is_bitwise, test_resume_is_bitwise).
Terminal window
ol start L0.6 # stubs checkpoint.py and tokens.py; safetensors.py is yours: ol start prints its contract diff
ol tests L0.6 # read the test catalog first: rung R0
ol check L0.6 # also reruns L0.0's smoke tests against your safetensors.py
ol check L0.6 --ref-deps # only if a dependency is not passing yet
ol diff L0.6 # after passing: your code against the reference

Your CLI’s train bigram --method autograd gains the token-stream form MS-L0 runs: --data shard.bin --max-steps N --batch B --seq-len T --ckpt-every K [--resume] reads windows through TokenStream, writes save_checkpoint(<out>/ckpt, ...) every K steps with extra = {"tokens_seen", "data_cursor", "lr", "config_sha256", "git_sha"} and the cursor’s generator state as rng_state, and with --resume restores model, optimizer, step, and cursor from load_checkpoint(<out>/ckpt). It evaluates the failpoint train/after-step after every step (the kata of craft.03 parses it) and ends with the final.json line MS-L0 compares byte for byte.


Three things in your system cannot happen yet. Your safetensors writer (L0.0) knows only F32, but SmolLM2’s weights (L7.9) are BF16 and quantized weights (M09.4, L8.5) are F8_E4M3 and U8. Your training runs keep their state in memory: MS-L0’s fourth step kills one after step 60, and today it would restart from zero, because nothing on disk says where it was, what AdamW’s moments were, or which window comes next. And the bigram trains on a byte string read whole into memory; the capstone’s TinyStories tokens are hundreds of megabytes in llm.c’s .bin format. This module writes every dtype, saves a training run so that a crash at any instant leaves a usable checkpoint, and reads token shards through a memory map with a cursor that resumes exactly.

SymbolMeaningType / shape
b31…b0b_{31} \dots b_0the 32 bits of a float32: sign, 8 exponent bits, 23 mantissa bits
BF16(x)\mathrm{BF16}(x)the top 16 bits of float32 xx after roundinguint16
s,e,ms, e, mF8_E4M3 fields: 1 sign bit, 4 exponent bits (bias 7), 3 mantissa bitsints
v(c)v(c)the value of an F8_E4M3 code ccfloat
TTseq_len, tokens per training window inputint
BBwindows per batchint
nkn_ktokens in shard kkint
ϕ\phithe phase of a pass, the offset of its first windowint in [0,min⁡(T,n−T))[0, \min(T, n - T))
oothe cursor offset: where the next window startsint
step- nnnnnnnnnnnna checkpoint directory, six digits of the optimizer steppath

Many dtypes, one layout. A safetensors file is unchanged from L0.0: an 8-byte header length, compact JSON, the tensor bytes back to back. What changes is the order when dtypes mix. The pinned library sorts tensors by dtype first, larger alignment first (F32, I32, BF16, F16, F8_E4M3, I8, U8 among the course dtypes), then by name as UTF-8 bytes (writer rule 1 of formats/safetensors.md). Name order alone gives the same file only when every tensor has one dtype, which is why L0.0’s goldens never saw the difference. A dtype comes from the array (float32 F32, float16 F16, int8 I8, uint8 U8, int32 I32) unless dtypes={name: ...} says otherwise; float64 is never converted silently, because a float64 checkpoint read as F32 by an engine is a different model.

BF16 is the top half of a float32. BF16 keeps float32’s sign and 8-bit exponent and only 7 of its 23 mantissa bits, so its bits are b31…b16b_{31} \dots b_{16} of the float32. Dropping the low 16 bits by truncation rounds every weight toward zero; the rule is round to nearest, ties to even, which M09.1’s f32_to_bf16_bits implements by adding 0x7FFF plus the lowest kept bit before shifting. Decoding is exact: shift the 16 bits back to the top of a float32. So the reader returns BF16 tensors as float32 arrays, and writing those values back as BF16 reproduces the file.

F8_E4M3 is a table of 256 codes. The “fn” variant (finite, NaN only) has bias 7 and no infinity:

eevalue
0 (subnormal)(−1)s⋅m8⋅2−6(-1)^s \cdot \frac{m}{8} \cdot 2^{-6}
1 to 15(−1)s⋅(1+m8)⋅2e−7(-1)^s \cdot (1 + \frac{m}{8}) \cdot 2^{e - 7}, except e=15,m=7e = 15, m = 7
15 with m=7m = 7NaN (0x7F, 0xFF)

The largest value is e=15,m=6e = 15, m = 6: 1.75⋅28=4481.75 \cdot 2^8 = 448. The smallest positive is 2−92^{-9} (code 0x01). Unlike IEEE formats, exponent 15 is not reserved for infinity, which is the habit to unlearn. To encode, round to the nearest code, ties to the even one (its last mantissa bit is 0; consecutive codes alternate that bit, across exponents too). A value beyond 448 has no code to round to: raise, so the missing scale is found (M09.4 adds the scale); NaN encodes as 0x7F.

Read the header without the data. read_header checks all five reader rules from the header and the file size alone (rule 5, the tiling, needs only the offsets and the size) and returns each tensor’s dtype and shape. It is what tells BF16 from F32 after load_safetensors decoded both to float32, and it is how a memory-mapping reader (L10.1) checks a file before trusting its offsets.

What a checkpoint holds. formats/checkpoint.md: model.safetensors (the state_dict, F32), optimizer.safetensors (each optimizer array named <parameter>.<state key>, for AdamW fc.weight.exp_avg and fc.weight.exp_avg_sq, for SGD fc.weight.momentum_buffer, with the rest of the optimizer’s state, its step count and hyperparameters, as JSON in the metadata key state), trainer_state.json (step, tokens seen, the data cursor, the generator state as 16 hex digits each so no JSON reader rounds a 64-bit integer, the learning rate, the config’s sha256, the git sha), config.json when there is one, and MANIFEST.json: every other file with its sha256 and size, sorted by name.

Writing atomically. A crash can stop a program between any two system calls, and until fsync returns, written bytes may exist only in the page cache. The protocol:

  1. write every file into step-<n>.tmp/ and fsync each;
  2. write MANIFEST.json last, fsync it, fsync the .tmp directory;
  3. rename the directory to step-<n> and fsync the parent (a rename is atomic: the directory appears whole or not at all);
  4. write LATEST.tmp with step-<n>\n, fsync, rename it to LATEST, fsync the parent;
  5. only then delete stale .tmp directories and, with keep, old steps.

So at every instant each step-<n> directory is complete and LATEST names one of them. Deleting old steps before step 4 would leave nothing to resume from if the save then failed.

Loading defensively. load_checkpoint(dir) starts from the directory LATEST names and checks it against its manifest: every listed file present with its size and sha256, and no file the manifest does not list. A directory that fails is skipped for the newest older one that passes. Directories newer than LATEST are ignored: they were never declared done. With no LATEST, the newest complete directory wins; with none complete, FileNotFoundError.

The token stream. A shard is a formats/tokens-bin.md file: a 1024-byte header of 256 little-endian int32 (20240520, version, nn, vocab size), then nn ids as uint16 (version 1) or uint32 (version 2). open_tokens validates the header and the file size and maps the ids with np.memmap, read-only, so a 100M-token shard costs nothing until a window is read. A window is T+1T + 1 consecutive ids: inputs are the first TT, targets the last TT (the next token of each input). Windows of one pass start TT apart, so each window’s last id is the next one’s first and every token after the phase is a target exactly once. A pass over shard kk starts at a random phase ϕ\phi = rng.below(min(T, n_k - T)) (the bound keeps one window in range on a short shard), so window boundaries move from pass to pass, and runs until the next window would pass the end; then the next shard (wrapping) starts a new pass with a new phase. The generator is used only at the start of a pass.

The cursor is the whole position. cursor() returns the shard index, the offset of the next window, and the generator’s state. Restoring all three into a stream over the same shards continues bit for bit; without the generator state the next pass draws a different phase and the resumed run diverges silently.

A BF16 file. The format page’s tensor w=[[1,2],[3,4]]w = [[1, 2], [3, 4]] with metadata {"format": "tinyllm"}, written as BF16.

  1. 1.01.0 is float32 0x3F800000; its top half is 0x3F80, stored little-endian as 80 3f. Likewise 2.0→2.0 \to 0x4000 (00 40), 3.0→3.0 \to 0x4040 (40 40), 4.0→4.0 \to 0x4080 (80 40). All four are exact: their low 16 bits are zero.
  2. The header {"__metadata__":{"format":"tinyllm"},"w":{"dtype":"BF16","shape":[2,2],"data_offsets":[0,8]}} is 93 bytes (the F32 example’s, with BF16 one byte longer and 8 one shorter than 16); three spaces pad it to N=96N = 96.
  3. The file: 60 00 00 00 00 00 00 00, the 96 header bytes, then 80 3f 00 40 40 40 80 40: 8+96+8=1128 + 96 + 8 = 112 bytes.

Rounding. 1+2−81 + 2^{-8} (float32 0x3F808000) sits exactly halfway between BF16 0x3F80 (11) and 0x3F81 (1+2−71 + 2^{-7}): the tie goes to the even code, 0x3F80. In F8_E4M3, 1.01.0 is e=7,m=0e = 7, m = 0, code 0x38, and 1.1251.125 is 0x39; 1.06251.0625 is halfway and becomes 0x38. 250250 lies between 240240 (0x77) and 256256 (0x78), above their midpoint 248248, so it becomes 0x78.

Token windows. One shard holding the ids 10,11,…,1910, 11, \dots, 19 (n=10n = 10), T=3T = 3, B=2B = 2, and phase draws that come out 1, then 0. Each draw is below(min(3, 10 - 3)) = below(3).

  1. Pass 1 starts at offset 1. Windows at offsets 1 and 4: ids [11,12,13,14][11, 12, 13, 14] and [14,15,16,17][14, 15, 16, 17]. Batch 1 is inputs [[11,12,13],[14,15,16]][[11, 12, 13], [14, 15, 16]], targets [[12,13,14],[15,16,17]][[12, 13, 14], [15, 16, 17]].
  2. The next window would start at 7 and need ids up to index 10, past the end. Pass 2 starts at phase 0: windows at 0 and 3, inputs [[10,11,12],[13,14,15]][[10, 11, 12], [13, 14, 15]].
  3. The cursor is now shard 0, offset 6 (the next window), plus the generator state.

The header. The ids [1,2,3][1, 2, 3] as version 1, vocabulary 256: 88 d8 34 01 (2024052020240520), 01 00 00 00, 03 00 00 00, 00 01 00 00, 1008 zero bytes, then 01 00 02 00 03 00: 1030 bytes.

A checkpoint. A Linear(2, 1) trained two AdamW steps and saved at step 7: step-000007/ holds model.safetensors (weight, bias), optimizer.safetensors (weight.exp_avg, weight.exp_avg_sq, bias.exp_avg, bias.exp_avg_sq), trainer_state.json, MANIFEST.json; LATEST holds step-000007 and a newline.

These are test_worked_example_bf16_bytes, test_f8_rounds_to_nearest_even, test_hand_example_token_windows, test_header_worked_example, and test_checkpoint_roundtrip.

# python/tinyllm/io/safetensors.py (taken over from L0.0; v0 calls keep working)
DTYPE_SIZES: dict[str, int]; E4M3_MAX = 448.0
def save_safetensors(path, tensors, meta, dtypes=None) -> None
def read_header(path) -> tuple[dict[str, tuple[str, tuple[int, ...]]], dict[str, str]]
def load_safetensors(path) -> tuple[dict[str, NDArray], dict[str, str]]
# python/tinyllm/io/checkpoint.py
class Checkpoint: path, step, model, opt, rng_state, extra
def step_name(step) -> str
def save_checkpoint(dir, model, opt, step, rng_state, extra, keep=None) -> str
def verify_step_dir(path) -> list[str]
def load_checkpoint(dir, step=None) -> Checkpoint
# python/tinyllm/io/tokens.py
def read_tokens_header(path) -> dict[str, int]; def open_tokens(path) -> NDArray
class TokenStream:
def __init__(self, shards, seq_len, batch, rng, vocab_size=None)
def next_batch(self) -> tuple[NDArray, NDArray]; def cursor(self) -> dict; def restore(self, cursor) -> None
TestKINDChecksWhy it matters downstream
test_worked_example_bf16_bytesunitsection 3’s 112-byte file, byte for byteyou and the test agree on BF16
test_matches_library_bytesgoldenfive library files, every dtype, written from float32 valuesany reader loads your files
test_raw_bits_are_written_as_isunituint16 BF16 bits and uint8 F8 codes go to disk untouchedre-saving HF or quantized weights
test_bf16_rounds_to_nearest_evenboundarythe ties of section 3 and a below-half casemixed precision (L11.1)
test_f8_e4m3_codesunitzero, 2−92^{-9}, 2−62^{-6}, 1, 448, NaN, −0-0, −448-448fp8 weights (M09.4)
test_f8_rounds_to_nearest_evenboundarysection 3’s F8 ties, a subnormal tie, −0-0the same codes as torch
test_f8_rejects_out_of_rangeboundary449 and infinity are errors; NaN is 0x7Fa missing scale fails loudly
test_dtype_order_then_nameunitF32, I32, U8 laid out in that order whatever the nameswriter rule 1
test_save_rejects_bad_dtypesboundaryfloat64, int64, unknown names, stray dtypes keys, BF16 from float16no silent conversion
test_load_decodes_every_dtypegoldenvalues, numpy types, shapes, writable arrays, signed zerosL7.9 reads BF16 weights
test_read_header_checks_without_readingunitdtypes and shapes; trailing bytes and a huge length rejectedmemory-mapped loading (L10.1)
test_load_rejects_unknown_dtypesboundaryF64, BOOL, bf16, a wrong byte countreader rule 4
test_raw_codes_roundtrippropertybits to values to bits gives the same file for every codedecode and encode agree
test_v0_calls_still_workregressionthe Pass 1 calls and the F32 fileyour CLI and engine still work
test_checkpoint_roundtripunitsection 3’s directory, LATEST, and every value loaded back bit for bitresuming
test_layout_matches_the_formatsconformancetrainer state keys and hex, sorted manifest with true hashes, tensor names, unlisted filesdur.09 and dur.12 read these files
test_config_and_defaultsunitconfig.json written and hashed; defaults for missing fieldsa minimal caller is still schema-valid
test_sgd_momentum_namesunitSGD state stored by parameter name and loaded backany optimizer checkpoints
test_resume_is_bitwiseproperty3 steps, save, load into a fresh model, 3 steps equals 6MS-L0 step 4
test_crash_at_every_write_keeps_a_valid_checkpointfaulta crash at each fsync or rename leaves complete directories and a valid LATEST; at least 8 fsyncsthe atomic protocol
test_load_skips_incomplete_and_corruptfaulta flipped byte, a missing manifest, a stale LATEST, nothing valida crash or bit rot cannot load garbage
test_newer_than_latest_is_ignoredboundaryLATEST governs; without it the newest complete winsthe rename order’s meaning
test_keep_and_stale_tmpunitkeep=2 leaves two steps; a stale .tmp is removeddisk use of long runs
test_rejects_bad_stateboundaryunknown extra keys, an even increment, a negative cursor, a short git sha, a negative step; nothing writtenschema-valid state or nothing
test_hand_example_token_windowsunitsection 3’s windows, cursor, and below(3) drawsyou and the test agree on the stream
test_header_worked_exampleunitthe 1030-byte file of tokens-bin.mdthe llm.c layout
test_version_2_reads_uint32unituint32 ids above 65535vocabularies above 64k
test_open_tokens_maps_the_fileunita read-only np.memmapshards larger than memory
test_rejects_bad_filesboundarywrong magic or version, size mismatch, out-of-vocabulary id, short filebad data fails before training
test_windows_tile_each_passpropertywindows TT apart, passes end where the next window would not fit, shards in orderevery token is a target once per pass
test_cursor_restore_is_bitwiseproperty7 batches, cursor, a new stream restored, 9 more equals one streamMS-L0 step 4
test_targets_are_inputs_shiftedpropertytargets[:, t] == inputs[:, t + 1], fresh int64 arraysnext-token prediction
test_stream_rejects_bad_argsboundaryno shards, a bare string, TT or B<1B < 1, too short, vocabulary mismatch; a shard with one windowerrors where the cause is visible
test_restore_rejects_bad_cursorboundaryshard or offset out of range; offset at the end starts a new passa cursor from another run
PitfallSymptomCaught by
1. laying a mixed-dtype file out by namesingle-dtype files match the library, mixed ones do nottest_dtype_order_then_name, test_matches_library_bytes (mutant s01)
2. truncating to BF16 or F8, or breaking ties upwardevery weight biased; bytes differ from torch’stest_bf16_rounds_to_nearest_even (mutant s02), test_f8_rounds_to_nearest_even (mutant s04)
3. saturating F8_E4M3 at 448out-of-range weights clip silently, the model degradestest_f8_rejects_out_of_range (mutant s07)
4. decoding F8_E4M3 subnormals as normals, or exponent 15 as NaNsmall weights come out wrong; 256 to 448 become NaNtest_f8_e4m3_codes (mutants s05, s06)
5. loose dtype checks: by kind only, or case-insensitive namesfloat64 slips in as F32; a file another reader rejects loads heretest_save_rejects_bad_dtypes (mutant s09), test_load_rejects_unknown_dtypes (mutant s11)
6. writing straight into step-<n>, skipping fsync, or replacing LATEST before the renamea crash leaves a half-written directory or a LATEST naming nothingtest_crash_at_every_write_keeps_a_valid_checkpoint (mutants s14, s15, s16)
7. a loader that trusts LATESTa corrupt newest checkpoint loads and the run trains on garbagetest_load_skips_incomplete_and_corrupt (mutant s17)
8. trusting the .bin headerversion 2 read as uint16, trailing bytes, ids beyond the vocabularytest_version_2_reads_uint32 (mutant s31), test_rejects_bad_files (mutants s33, s34)
9. windows T+1T + 1 apart, or targets equal to inputsboundary tokens are never targets; the model learns the identitytest_hand_example_token_windows (mutant s25), test_targets_are_inputs_shifted (mutant s26)
10. a cursor without the generator statea resumed run draws a different phase at its next pass and divergestest_cursor_restore_is_bitwise (mutant s29)
11. the phase from below(T) on a short sharda window past the end of the shardtest_stream_rejects_bad_args (mutant s24)
12. a cursor naming the window just reada resumed run repeats one windowtest_hand_example_token_windows (mutant s30)

| Forward | L2.2 | Registered call site uses this module. | | Forward | L3.6 | Registered call site uses this module. | | Forward | L4.1 | Registered call site uses this module. | | Forward | L5.5 | Registered call site uses this module. | | Forward | L6.1 | Registered call site uses this module. | | Forward | L6.2 | Registered call site uses this module. | | Forward | L6.3 | Registered call site uses this module. | | Forward | L6.5 | Registered call site uses this module. | | Forward | L6.7 | Registered call site uses this module. | | Forward | L7.9 | Registered call site uses this module. | | Forward | dur.09 | Registered call site uses this module. |

DirectionModuleHow it uses this
BackM09.1f32_to_bf16_bits and bf16_bits_to_f32 for BF16
BackL0.1the tensors of the model the checkpoint tests train
BackL0.2F.mean in the checkpoint tests’ loss
BackL0.4a checkpoint is state_dict(); resume is load_state_dict()
BackM10.2SGD’s momentum buffers, saved by parameter name
BackM10.3AdamW’s moments and step count, saved and restored
BackM06.3PCG32.below draws each pass’s phase; state() and set_state() move with the cursor
BackL0.5L0.0’s suite runs as your regression for safetensors.py and also covers bigram.py, which L0.5 owns
ForwardL10.0your engine reads the model.safetensors this writer produces
ForwardL2.1the n-gram model writes and reads its count tables through save_safetensors and load_safetensors

Later passes build on it without changing it: L2.2 reads TokenStream windows, L7.9 loads BF16 SmolLM2 weights, L10.1 memory-maps the same layout from Rust, and the capstone trainer and dur.09 resume from these checkpoints.

Your pieceProduction equivalentWhat it addsWhere to look
save_safetensorsHugging Face safetensorszero-copy memory-mapped loading, lazy slicing of one tensor, sharded files with an indexsafetensors/src/tensor.rs
F8_E4M3 tableml_dtypes, torch.float8_e4m3fnE5M2, scaled matmuls on Hopper, MX block-scaled formatsml_dtypes/_src/float8.h
atomic checkpointsPyTorch Distributed Checkpoint, Orbaxsharded saves from many ranks, async writes off the training thread, a commit marker per savetorch/distributed/checkpoint/, orbax/checkpoint/
TokenStreamllm.c DataLoader, nanotron, MosaicML streamingsharding across ranks, shuffled shard order, deterministic resumption across world sizesllm.c llmc/dataloader.h, streaming/base/dataset.py