Skip to content

Model Loading: Hub to GPU

How a checkpoint becomes a running model, and how to debug it when it does not. Every “my custom fine-tune won’t load / is slow to start / gives garbage” ticket at an inference cloud reduces to one of four layers: the repo (files and metadata), the bytes (format, dtype, layout), the loader (key mapping, sharding, quantization), or the tokenizer (template and special tokens). This chapter works through each layer as math first, then the code that checks it.

Parent topic: LLM Systems & Inference. Formats consumed here are explained in Quantization: Math → Code; the runtimes that do the loading are mapped in Inference Frameworks and LLM Serving Platforms. Runnable examples: code/.

  • A model is config + weights + tokenizer + generation defaults, and each lives in a different file. Garbage output is more often a tokenizer or config bug than a weights bug
  • safetensors is a JSON header plus a flat byte buffer. You can inspect any checkpoint, local or remote, by reading the first 8 + N bytes; you never need to load it to know its dtypes, shapes, and size
  • from_pretrained is a key-matching problem. Every load is “map checkpoint names to module parameters, then fill them”; prefix drift, tied weights, and fused projections cause most load failures
  • Cold start is a bandwidth equation: T ≈ bytes / min(B_net, B_disk, B_pcie) plus engine init. Production loads from local NVMe or parallel object-store reads, never from the Hub at request time
  • Chat templates fail silently. A wrong template still produces fluent text; only diffing token ids against training proves correctness
  • Run code/safetensors_inspect.py --url against three Hub models (one tied-embedding, one FP8, one MXFP4) and explain every tensor name you see
  • Run code/config_explain.py on the samples, then hand-derive the Llama 3.1 8B parameter count until you get 8,030,261,248 exactly (it matches metadata.total_size / 2 in the official index)
  • Break a load on purpose: rename keys with a module. prefix, drop lm_head.weight, delete rope_scaling, swap the chat template. Record the symptom of each; that table is your on-call runbook
  • Read one real loader end to end: vLLM’s llama.py hf_to_vllm_mapper and AutoWeightsLoader

Loading is a bijection problem under a byte budget. A checkpoint is a dictionary {name → (dtype, shape, bytes)}; a model is a dictionary {name → parameter slot}. Loading succeeds when there is a renaming function f and a set of tensor operations (cast, transpose, concatenate, slice per rank, dequantize) that map one onto the other with nothing missing and nothing left over. It is fast when those bytes stream at the slowest link’s bandwidth with no extra copies. It is correct only when the tokenizer feeds the model the same token ids it saw in training. Everything below is a specialization of those three statements.

A Hub model repo is a git repository with large files stored out of band (Xet, the successor to Git LFS: content-defined chunks, deduplicated across repos and revisions).

FileWhat it decidesFields an FDE checks first
config.jsonModel class and shapearchitectures, model_type, hidden_size, num_attention_heads, num_key_value_heads, head_dim, rope_theta, rope_scaling, max_position_embeddings, tie_word_embeddings, vocab_size, torch_dtype/dtype, quantization_config, auto_map
generation_config.jsonDefault decodingeos_token_id (often a list), bos_token_id, pad_token_id, temperature, top_p, do_sample
tokenizer.jsonThe tokenizer itself (fast, tokenizers crate)model.type (BPE/Unigram), added_tokens, pre_tokenizer, post_processor (adds BOS?)
tokenizer_config.json / chat_template.jinjaSpecial tokens and the chat templatebos_token, eos_token, pad_token, added_tokens_decoder, chat_template
model.safetensors or model-0000k-of-0000n.safetensorsWeightsdtype per tensor, presence of lm_head.weight
model.safetensors.index.jsonShard mapmetadata.total_size (bytes), weight_map (tensor name → shard file)
README.md (model card)License, base model, intended templateYAML front matter: base_model, license, pipeline_tag
LICENSE, gatingLegal accessGated repos ("gated": "manual") need an accepted license plus HF_TOKEN
adapter_config.json + adapter_model.safetensorsA LoRA adapter, not a full modelbase_model_name_or_path, r, lora_alpha, target_modules

Key ideas:

  • model_type is the dispatch key. AutoModelForCausalLM looks it up in a registry (llama → LlamaForCausalLM). If the installed library does not know the type, loading fails before a byte of weights is read
  • auto_map means custom code. It points at modeling_*.py files in the repo, which only run with trust_remote_code=True. That executes arbitrary Python from the repo; pin a commit revision if you must use it
  • Revisions are commits. main moves; a 40-char sha does not. A customer saying “it worked yesterday” on main often means the repo changed. Resolve once and pin: revision="0e9e39f..."
  • Config dtype is a hint, not a guarantee. torch_dtype (renamed dtype in transformers v5, old key still read) records the training dtype; the safetensors header records what is actually stored. When they disagree, trust the header

The real Llama 3.1 8B Instruct config (in code/samples/) shows three classic traps in eight lines: eos_token_id is a list [128001, 128008, 128009], rope_scaling.rope_type = "llama3" (older libraries reject it), and num_key_value_heads = 8 < num_attention_heads = 32 (GQA, which matters for TP sharding).

offset 0 8 8+N EOF
┌────────┬────────────────────┬──────────────────────────────────────┐
│ N (u64 │ JSON header, UTF-8 │ byte buffer: tensors back to back, │
│ LE) │ padded with ' ' │ row-major, little-endian, no holes │
└────────┴────────────────────┴──────────────────────────────────────┘
header = {"model.norm.weight": {"dtype":"BF16","shape":[576],"data_offsets":[b,e]},
"__metadata__": {"format":"pt"}} # values must be strings
tensor bytes live at file[8 + N + b : 8 + N + e], and e - b = prod(shape) * bits(dtype) / 8

Math → code (code/safetensors_inspect.py):

(n,) = struct.unpack("<Q", f.read(8)) # header length
header = json.loads(f.read(n)) # starts with '{'
b, e = header[name]["data_offsets"]
f.seek(8 + n + b); raw = f.read(e - b) # zero parsing of the payload

Key ideas:

  • Zero-copy and mmap: because offsets are known up front, a loader can mmap the file and hand out tensor views with no deserialization. Pages fault in lazily on first touch, which is why a “0.3 s load” on a warm page cache becomes 30 s on a cold node. Measure cold with sync; echo 3 > /proc/sys/vm/drop_caches
  • Remote inspection: two HTTP Range requests read the header of a 100 GB checkpoint. safetensors_inspect.py --url does this, so you can verify a customer’s dtypes and key names before pulling a byte of weights
  • Validation: the spec forbids holes and duplicate keys; END - BEGIN must equal prod(shape) × bits / 8. A mismatch means a corrupt or hand-edited file
  • Sharding: save_pretrained writes shards plus model.safetensors.index.json. A missing shard is a common partial-upload failure; compare weight_map against the shards present

pytorch_model.bin is a zip of Python pickles. Unpickling executes the __reduce__ callables the file names, so loading an untrusted .bin is remote code execution (Python docs warning). PyTorch 2.6 made torch.load(weights_only=True) the default, which restricts unpickling to tensors and primitives, but inference platforms should still refuse .bin from customers and convert to safetensors in a sandbox. safetensors cannot execute code: it is JSON plus raw bytes.

FormatStructureWho loads itPortability
safetensorsJSON header + flat buffertransformers, vLLM, SGLang, TGI, candle, MLXAny framework, any hardware
PyTorch .bin/.ptZip of picklesLegacy PyTorchUnsafe, Python-only
GGUFMagic GGUF, version, KV metadata (arch, tokenizer, tokenizer.chat_template), tensor infos, aligned datallama.cpp, Ollama, LM Studio, mistral.rs, candleOne file holds weights and tokenizer and template; k-quant block types
ONNXProtobuf graph + initializers (external data > 2 GB)ONNX Runtime, DirectML, edgeGraph is frozen; dynamic shapes need care
TensorRT engineSerialized, compiled kernelsTensorRT / TensorRT-LLMTied to GPU arch, TRT version, and build-time max shapes; rebuild per SKU
Orbax checkpointDirectory of array shards (TensorStore/OCDBT)JAX, Flax NNX, MaxTextSharding-aware restore onto a device mesh
Keras .weights.h5 / preset dirHDF5 or preset folderKeras 3, KerasHubBackend-agnostic (JAX, PyTorch, TF)
safetensors dtypeBits (sign/exp/mantissa)Max finiteWhere it appears
F321/8/233.4e38Norms in some checkpoints, optimizer state
F161/5/1065504Older checkpoints, AWQ/GPTQ scales
BF161/8/73.4e38Default for modern LLM weights
F8_E4M31/4/3448FP8 weights (DeepSeek-V3, compressed-tensors, ModelOpt)
F8_E5M21/5/257344Gradients, some KV caches
F8_E8M00/8/0 (power of two)2^127MX block scales
F4 (E2M1)1/2/16MXFP4 / NVFP4 elements
U8, I32containersn/aPacked low-bit weights: GPTQ/AWQ qweight (8 × int4 per I32), gpt-oss MXFP4 *_blocks (2 × fp4 per U8)

MXFP4 (OCP MX): blocks of 32 E2M1 values share one E8M0 scale, about 4.25 bits per weight. NVFP4 (Blackwell): blocks of 16 E2M1 values share an FP8 E4M3 scale plus a per-tensor FP32 scale. In a header these look like a U8 or F4 tensor plus a *_scales tensor, so element count ≠ parameter count for quantized checkpoints. DeepSeek-V3 stores F8_E4M3 weights with a weight_scale_inv tensor per 128×128 block. The quantizers themselves are derived in Quantization.

from huggingface_hub import hf_hub_download, snapshot_download
cfg = hf_hub_download("Qwen/Qwen3-30B-A3B", "config.json") # one file
path = snapshot_download( # a whole revision
"meta-llama/Llama-3.1-8B-Instruct",
revision="0e9e39f249a16976918f6564b8830bc894c89659", # pin a commit
allow_patterns=["*.json", "*.safetensors", "tokenizer*"], # skip .bin, GGUF, original/
)
Terminal window
hf auth login # or export HF_TOKEN=hf_...
hf download meta-llama/Llama-3.1-8B-Instruct --include "*.safetensors" --include "*.json" \
--revision 0e9e39f249a16976918f6564b8830bc894c89659 --local-dir /mnt/nvme/llama31-8b
hf cache ls # what is cached, how big
hf cache verify meta-llama/Llama-3.2-1B-Instruct # checksum against the Hub
HF_HUB_OFFLINE=1 python serve.py # never touch the network

The CLI is hf; huggingface-cli is the deprecated name (CLI guide).

Cache layout ($HF_HOME/hub, default ~/.cache/huggingface/hub, override with HF_HUB_CACHE):

models--meta-llama--Llama-3.1-8B-Instruct/
├── blobs/<sha256 or etag> # actual bytes, content-addressed
├── refs/main # text file: commit sha that "main" resolved to
├── snapshots/<commit-sha>/
│ ├── config.json -> ../../blobs/…
│ └── model-00001-of-00004.safetensors -> ../../blobs/…
└── trees/<commit-sha>.json # cached file list for that commit

Key ideas:

  • Symlinks dedupe revisions: two revisions sharing a shard point at one blob. On filesystems without symlinks (some network mounts, Windows without developer mode) files are copied, and disk use multiplies. Recent huggingface_hub also shares Xet blobs across repos, so fine-tunes that reuse base shards cost no extra disk
  • --local-dir writes plain files instead of the cache layout. Use it when baking weights into an image or a volume
  • Speed: hf_xet (installed with huggingface_hub) replaced hf_transfer. Set HF_XET_HIGH_PERFORMANCE=1 on big-NIC machines to raise concurrency. Without HF_TOKEN, requests are rate-limited
  • Offline: HF_HUB_OFFLINE=1 makes every call resolve from cache; with a pinned sha, a fully cached load makes zero network calls. An incomplete snapshot raises IncompleteSnapshotError instead of silently returning a partial folder
  • Production rule: the Hub is a distribution channel, not a serving dependency. Mirror pinned revisions once into your own object store (s3://models/<org>/<name>/<sha>/), checksum them, and load from there or from local NVMe. Reasons: rate limits and outages, gated-token sprawl, revision drift, and egress cost across thousands of replicas
from_pretrained(repo_or_dir, dtype="auto", device_map="auto")
1. resolve files snapshot / local dir; read config.json (+ quantization_config)
2. pick class config.model_type → AutoModel mapping → LlamaForCausalLM
(auto_map + trust_remote_code → repo's modeling_*.py instead)
3. skeleton on meta build modules on the "meta" device: shapes, no memory
4. plan placement device_map="auto" (accelerate): fill GPU 0..n, then CPU, then disk
5. stream state dict for each shard: rename keys → convert (fuse/split/dequant) → shard (TP)
→ cast to dtype → materialize on target device (4 threads by default)
6. finalize tie weights (lm_head ← embed_tokens if tie_word_embeddings),
init any missing params, report missing / unexpected / mismatched keys
7. generation config load generation_config.json → model.generation_config

In transformers v5 the argument is dtype (torch_dtype is the legacy name), and the default now follows the checkpoint’s config dtype instead of upcasting to FP32. Step 5 is the dynamic weight loader: WeightRenaming rules (e.g. LayerNorm.gamma → LayerNorm.weight) and WeightConverter rules (e.g. stack Mixtral’s per-expert w1/w3 into one experts.gate_up_proj with MergeModulelist + Concatenate) run as tensors stream in, and they are reversible so save_pretrained writes the original layout back. Peak memory ≈ model size plus the largest merge (one MoE layer’s experts).

Key ideas:

  • Meta device removes the “two copies” problem: without it you allocate random weights and then the checkpoint, doubling peak RAM
  • Read the load report. “Some weights were not initialized from the checkpoint” with lm_head.weight listed is not a warning; it means a random output head
  • trust_remote_code executes repo Python at load time. Inference platforms run it, if at all, in a sandbox at conversion time, never in the serving process. Note that custom-code models skip transformers’ built-in conversion mappings
SymptomRoot causeCheckFix
“weights not initialized” for every layer; unexpected keys start with module., _orig_mod., base_model.model.Saved from a DDP / torch.compile / PEFT wrapperDiff header keys against model.state_dict().keys()Strip the prefix; save from the unwrapped model; merge_and_unload() for PEFT
Output is random text from token 1lm_head.weight missing and tie_word_embeddings=false (or present but tied flag wrong)Is lm_head.weight in weight_map?Set tie_word_embeddings to match the checkpoint
Fine at 4K tokens, degrades past 8Krope_scaling dropped or rewritten by a fine-tune toolkit; or old library ignores rope_type: llama3/yarnDiff rope_scaling and rope_theta against the base configRestore the base values; upgrade the library
Fluent but worse than base, ignores system promptChat template differs from trainingRender both templates, diff token ids (§7)Ship the training template in tokenizer_config.json / chat_template.jinja
Never stops; prints assistant turns foreverEOS mismatch: instruct turn ends with <|eot_id|> (128009) but only <|end_of_text|> (128001) is an EOSgeneration_config.eos_token_id vs template’s turn terminatorMake eos_token_id a list containing every terminator
NaN / inf or garbage only in FP16BF16-trained activations exceed 65504Same prompt in BF16 is fineServe BF16 (or FP8 with scales), never FP16, for BF16-trained models
size mismatch for embed_tokens: [128264, 4096] vs [128256, 4096]Tokens added in fine-tune, or vocab padded to a multiple of 64len(tokenizer) vs config.vocab_size vs header shaperesize_token_embeddings(len(tok)) before saving; set vocab_size to the stored rows
model type X not recognizedLibrary older than the architecture, or custom codetransformers.__version__; auto_map in configUpgrade, or pinned-revision trust_remote_code in a sandbox
Quality collapse on a quantized fine-tunequantization_config stale (fine-tuned in BF16, config still says GPTQ), or ignore list missing lm_headHeader dtypes vs quantization_configRequantize from the merged BF16 weights

vLLM and SGLang keep their own model registries (architecture string → engine-native implementation with fused kernels) and their own loaders; they read the same config.json and safetensors but never call transformers’ from_pretrained for supported architectures. An unsupported architectures entry falls back to the Transformers backend (slower) or fails.

Engines fuse q_proj, k_proj, v_proj into one qkv_proj GEMM and gate_proj, up_proj into one gate_up_proj. vLLM’s Llama declares the mapping (source):

hf_to_vllm_mapper = WeightsMapper(
orig_to_new_stacked={
# weight_name: (param_name, shard_id)
".q_proj": (".qkv_proj", "q"),
".k_proj": (".qkv_proj", "k"),
".v_proj": (".qkv_proj", "v"),
".gate_proj": (".gate_up_proj", 0),
".up_proj": (".gate_up_proj", 1),
}
)

Each parameter carries a weight_loader(param, loaded_weight, shard_id) that knows where in the fused tensor, and which rows for this rank, the incoming tensor goes. A customer checkpoint that already fused qkv_proj (common from custom training code) misses this mapping: the fix is to split it back to HF names, not to patch the engine.

PyTorch stores nn.Linear weights as W ∈ ℝ^{out × in}. With t ranks:

ColumnParallel (q, k, v, gate, up): W_r = W[r·out/t : (r+1)·out/t, :] output split, no comm
RowParallel (o_proj, down_proj): W_r = W[:, r·in/t : (r+1)·in/t] partial sums → all-reduce
Vocab-parallel embedding / lm_head: rows r·V'/t … where V' = V padded up to a multiple of 64 (vLLM)

For GQA with H query heads and KV key/value heads, rank r gets query heads [r·H/t, (r+1)·H/t) and KV heads [r·KV/t, (r+1)·KV/t). Requirements: H mod t = 0; if KV < t, each KV head is replicated across t/KV ranks (so t mod KV = 0). Llama 3.1 8B (H=32, KV=8) shards cleanly to t ∈ {1,2,4,8}; Qwen3-30B-A3B (KV=4) at t=8 replicates each KV head twice.

The fused gate_up_proj = [G; U] must be sliced per sub-matrix: rank r gets [G_r; U_r], not the contiguous r-th chunk of the fused tensor (which would hand rank 0 all of G). Transformers’ TP plan calls this packed_colwise. Getting it wrong loads without error and produces garbage, which is why fused checkpoints are dangerous.

Engines read config.json → quantization_config.quant_method and swap in the matching linear method and kernel:

quant_methodStored tensorsKernel path (vLLM)
awqqweight (I32, 8×int4), qzeros, scales (F16), group 128AWQ Marlin on Ampere+
gptqqweight, qzeros, scales, g_idx (act-order)GPTQ Marlin / Machete
fp8F8_E4M3 weight + weight_scale_inv (block 128×128) or per-tensor scaleCUTLASS FP8 GEMM (Hopper+, Ada)
compressed-tensorsPer-group config_groups (W4A16, W8A8, FP8, NVFP4) + ignore listMarlin / CUTLASS by scheme (llm-compressor output)
modeloptNVIDIA ModelOpt FP8 / NVFP4Blackwell FP4 tensor cores
mxfp4gpt-oss experts as *_blocks (U8) + *_scales (E8M0)MXFP4 MoE kernels

--quantization on the CLI overrides detection; usually you should not pass it. The common ticket: “quantized fine-tune is slow” because the GPU lacks the fast kernel (FP8 on Ampere falls back to weight-only Marlin), not because loading failed.

An adapter is adapter_config.json (r, lora_alpha, target_modules, base_model_name_or_path) plus adapter_model.safetensors with keys like base_model.model.model.layers.0.self_attn.q_proj.lora_A.weight. The math:

W' = W + (α / r) · B A A ∈ ℝ^{r × in}, B ∈ ℝ^{out × r} (rsLoRA: α / √r)
extra params per adapted matrix = r · (in + out) e.g. r=16, 4096×4096 → 131,072 (0.8 %)
Merge (merge_and_unload(), then serve as a full model)Serve unmerged (vllm serve base --enable-lora --lora-modules name=path)
LatencyZero overheadExtra 2·r·(in+out) FLOPs/token per adapted matrix; batched SGMV/Punica kernels
MemoryFull copy per fine-tuneOne base, many adapters (multi-tenant)
Cold startFull model loadAdapter is MBs: seconds
GotchasMerging into a quantized base needs dequant → merge → requantr > --max-lora-rank (default 16) rejects; adapters that also train embed_tokens/lm_head (added tokens) may not be supported; base revision must match training

Rule: serve unmerged for many low-traffic fine-tunes; merge for one high-traffic fine-tune where every microsecond of TPOT matters.

T_cold = T_schedule + T_fetch + T_read + T_h2d + T_init
T_fetch ≈ S / B_net (only if weights are not already on the node)
T_read + T_h2d ≈ S / min(B_disk, B_pcie) when streamed and overlapped, else their sum
T_init = CUDA context + KV-cache profiling + CUDA graph capture (+ torch.compile)
S = params × bytes/param (per rank: S / t if checkpoints are pre-sharded)

Worked examples (BF16; 8B = 16.06 GB from the Llama 3.1 index, 70B = 70.55 B params = 141.1 GB):

PathEffective bandwidth8B (16.1 GB)70B (141 GB)
Hub or S3, single HTTP stream~0.2 GB/s80 s12 min
Object store, parallel range reads, 25 GbE~3 GB/s5.4 s47 s
Object store, parallel, 100 GbE~10 GB/s1.6 s14 s
Local NVMe Gen4 (one drive)~6.5 GB/s2.5 s22 s
NVMe RAID-0 ×4 or page cache~25 GB/s0.6 s5.6 s
PCIe Gen5 x16 host → one GPU~50 GB/s0.3 s2.8 s (÷ t with TP, each GPU has its own link)

Two lessons fall out. The network and disk terms dominate by 10-100×, so cache placement beats loader tuning. And once weights are local, T_init (graph capture, compile, often 20-60 s for large models) becomes the floor, which is why the last tricks below snapshot GPU state instead of reloading.

TechniqueAttacksMechanism
safetensors mmapCopiesLazy page-in, no deserialization; worst case on network filesystems (random 4 KB faults)
fastsafetensors (--load-format fastsafetensors)T_read + T_h2dBatched reads and GPUDirect Storage straight into GPU memory
Run:ai Model Streamer (--load-format runai_streamer)T_fetchConcurrent reads from S3/GCS/Azure/local into a CPU buffer, overlapped with H2D; --model-loader-extra-config '{"concurrency":16}'
tensorizer (--load-format tensorizer)T_fetchSerialize once, stream from HTTP/S3 at line rate
Pre-sharded state (sharded_state)T_read per rankSave each TP rank’s slice once; each rank reads S/t with no slicing
Local NVMe cache / warm poolT_fetchDaemonset pre-pulls pinned revisions; scale from warm nodes
GPU snapshotting (cuda-checkpoint + CRIU, platform GPU memory snapshots)T_initRestore a process with weights and CUDA graphs already resident

code/config_explain.py prints this estimate for any config.json.

FamilyAlgorithmExamplesTell-tale
BPE (Sennrich et al.)Greedy merges of frequent pairs from a merges tableGPT-2 lineagemerges in tokenizer.json
Byte-level BPEBPE over bytes, so no unknown tokensGPT-2, Llama 3 (128,256 vocab), QwenĠ marks a leading space
SentencePiece (repo)Unigram LM or BPE over raw text, whitespace as ▁Llama 2, Gemma, T5▁ prefix; tokenizer.model file
tiktoken (repo)Byte-level BPE with regex pre-split, fast Rust coreOpenAI cl100k_base, o200k_base; gpt-oss o200k_harmonyRanks file, no merges table

A chat template is a Jinja program stored in tokenizer_config.json (chat_template) or chat_template.jinja. It turns [{"role","content"}] into the exact string the model was trained on:

SmolLM2: <|im_start|>user\nHi<|im_end|>\n<|im_start|>assistant\n
Llama 3 (simplified; 3.1 also injects a dated system header):
<|begin_of_text|><|start_header_id|>user<|end_header_id|>\n\nHi<|eot_id|><|start_header_id|>assistant<|end_header_id|>\n\n

Why mismatches are silent: the model still sees valid tokens, so output is fluent. It is just conditioned on an out-of-distribution prefix: system prompts get ignored, tool calls malformed, refusals spike, evals drop a few points. Nothing errors.

How to verify (do this on every custom model onboarding):

ids_serving = tok.apply_chat_template(msgs, add_generation_prompt=True) # what the engine sends
ids_training = tok(training_formatter(msgs), add_special_tokens=False).input_ids # what SFT saw
assert ids_serving == ids_training, first_divergence(ids_serving, ids_training)

Common divergences: double BOS (the template already emits <|begin_of_text|>, then the rendered string is tokenized again with add_special_tokens=True, so the post-processor adds a second one), missing add_generation_prompt, a fine-tune that added special tokens (added_tokens) the serving tokenizer lacks, trailing whitespace or newline differences, and a template that drops the system role. vLLM uses the repo template unless --chat-template overrides it, and --generation-config auto pulls sampling defaults from generation_config.json, so both files are part of the deployed contract.

transformers v5 dropped its TensorFlow and Flax backends, so the JAX world loads through its own stack:

# Flax NNX + Orbax: build an abstract model (shapes only), then restore into it
abstract = nnx.eval_shape(lambda: Transformer(cfg, rngs=nnx.Rngs(0)))
graphdef, abstract_state = nnx.split(abstract)
state = ocp.StandardCheckpointer().restore(ckpt_dir / "state", abstract_state)
model = nnx.merge(graphdef, state)

nnx.eval_shape plays the role of PyTorch’s meta device, and Orbax restores each array directly onto its target sharding across a device mesh (no host-side gather), which is how MaxText and TPU serving load. Converting HF safetensors into JAX means the same key renaming plus transposes (nn.Linear stores out × in; Flax Dense kernels store in × out).

Keras 3 / KerasHub presets are backend-agnostic: keras_hub.models.CausalLM.from_preset("hf://<org>/<repo>", dtype="bfloat16") converts supported HF architectures (Llama 3, Gemma, Mistral, Mixtral, Qwen, gpt-oss, and others in utils/transformers/) on load and runs them on JAX, PyTorch, or TensorFlow.


Customer saysFirst commandMost likely cause
“Won’t load”safetensors_inspect.py <index.json> and diff keys vs base modelPrefix drift, missing shard, fused or renamed projections, unknown model_type
“Loads but OOMs”config_explain.py config.jsonWeights + KV cache at max context exceed GPU; dtype upcast to FP32
“Slow to start”Time each term of §6 separatelyPulling from the Hub per replica; network filesystem with mmap; no warm cache
“Gives garbage”Same prompt in transformers BF16 vs engineMissing lm_head, FP16 overflow, TP slicing of fused tensors, quant config mismatch
“Worse than the base model”Diff templated token ids vs trainingChat template, double BOS, EOS list, sampling defaults
“Never stops”Print generation_config.eos_token_idTurn terminator not in EOS list
ConceptConnected TrackApplication
Transformer blocks, GQA, MoE, RoPEDeep LearningWhat each tensor name means
FP8/MXFP4/NVFP4, AWQ, GPTQQuantizationReading pre-quantized checkpoints
Fine-tuning, LoRA, SFT data formattingTraining & Post-TrainingWhere adapters and templates come from
Image baking, daemonsets, warm poolsCloud NativeGetting weights onto nodes
Load-time and TTFT metricsObservabilityMeasuring each cold-start term
CompanyHow This AppearsDifficulty
Fireworks / Together / BasetenCustom-model onboarding, LoRA multi-tenancy, cold-start SLAs; FDEs debug customer checkpoints dailyExpert
Hugging FaceHub storage (Xet), safetensors, transformers loader, TGI/Inference EndpointsExpert
NVIDIAModelOpt FP8/NVFP4 checkpoints, TensorRT-LLM engine builds, GPUDirect StorageExpert
Modal / Replicate / RunPodScale-to-zero, GPU memory snapshots, weight cachingExpert
Anyscale / DatabricksvLLM-based serving, model registries, object-store streamingExpert
GoogleOrbax, MaxText, KerasHub presets, TPU servingExpert