Training and Inference Systems

Michael BrenndoerferAugust 9, 202656 min read

Part of World Models Handbook

Distributed training, checkpointing, and stateful rollout serving for world models, with latency, memory, and energy budgets for deployment.

Choose your expertise level to adjust how many terms are explained. Beginners see more tooltips, experts see fewer to maintain reading flow. Hover over underlined terms for instant definitions.

Article links

Make inline references clickable

Training and Inference Systems

A robot controller has 20 milliseconds to choose its next action. If its world-model query spends 25 milliseconds waiting in a queue, a fast forward pass cannot recover the missed deadline. That is an illustrative budget, not a benchmark of a particular robot, but it gives us a concrete systems problem: accurate predictions must arrive within the time and resource limits of their consumer. Up to this point we have talked about what world models are (Part I), how they represent state (Part III and Part IV), which architectures realize that representation (Part V), how they are learned (Part VI), and how they are used to plan and act (Part VII and Part VIII). Here we connect those ideas to training infrastructure, recoverable experiments and rollout serving. We will separate analytical cost models from local CPU measurements rather than infer deployment performance from model names.

The distinction between training and inference loops matters. Training advances through data and gradient updates. Autoregressive inference advances through predicted time steps, sometimes in lockstep with physical equipment. Both depend on numerical quality, but also on the machinery delivering the right inputs, state and outputs within a budget. A late accurate prediction and an inaccurate timely prediction are different failures; either can make a controller miss its requirements. Instrumenting both quality and timing lets us tell them apart.

A world-model training and inference stack includes the model's predict() call and the infrastructure around it. We will examine four components that need to work together:

  • Training infrastructure that distributes gradient computation when needed, preserves optimizer state across updates, and saves snapshots for recovery.
  • Data, checkpoint, and experiment infrastructure that move trajectory data and manifests between storage and compute, and record the artifacts needed to compare or reconstruct runs.
  • Rollout serving that exposes the learned dynamics as a stateful, low-latency service, with the caching and batching strategies specific to autoregressive latent rollouts.
  • Latency, memory, and energy optimization that turns an algorithm into something a robot, a browser, or a data center can run within its power and time budget.

These components are not independent. Their boundaries can introduce failures that component-only tests miss. Data availability can enter the training-step critical path. Checkpoint overhead and retention determine the recovery points available after a failure. The rollout server's cache footprint constrains concurrent sessions and feasible batches, which in turn affect response time. After changing one component, measure the whole path again: the limiting resource may have moved.

Consider three failures. A rollout service substitutes a default when overloaded, changing the training data without raising a model exception. A checkpoint mixes tensors from before and after an optimizer update, changing the resumed computation. A cache key omits episode identity and returns another episode's state. These failures can produce plausible numbers rather than an obvious crash. Explicit state identities, checksums, invariants and continuation tests can expose them; silence is a risk to test for, not an inevitable property of every failure.

This chapter assumes the model families from Part V, learning methods from Part VI, and decision-centric uses from Part VIII. We treat the model as a potentially expensive, stateful function and ask how to execute it correctly and within budget. Many techniques are shared with language and vision models. The workloads considered here can add long autoregressive rollouts, sensor timing and planning consumers that amplify some prediction errors. Those features are not universal to every world model, and a contractive transition can attenuate rather than amplify perturbations.

Let's begin with the training side.

Distributed Training and Sequence Parallelism

Distributed training is the collection of techniques that allow one logical model to be trained using the memory and compute of many accelerators. It exists for two reasons that should be kept separate in your head, because they call for different remedies and are often confused with one another:

  • Memory capacity: Parameters, gradients, optimizer state, activations and the input batch must fit. Two fp32 Adam moment tensors for a billion parameters occupy 8 GB decimal, excluding step counters and any master-weight copy. That is not the entire training footprint. Dtypes, optimizer variants and sharding determine what each rank actually stores.
  • Throughput and wall-clock: Even when everything fits, dividing work can reduce completion time. A hypothetical day-long run that finishes in one hour across thirty-two workers has a 24-fold speedup, not ideal 32-fold scaling. Throughput affects experiment iteration and cost; under a fixed deadline or budget it can also determine whether a run is feasible.

The distinction matters because the two problems have different fixes. You should identify which tensors cause the memory limit. Tensor and pipeline parallelism partition model computation. Sharded data parallelism can instead partition optimizer states, gradients and parameters while retaining data-parallel work assignment, as described in ZeRO. Activation checkpointing trades recomputation for activation memory. Adding fully replicated data-parallel workers does not reduce the replicated parameter or optimizer footprint.

Trajectory length creates another possible memory limit. As an illustrative input, a five-minute video at 30 Hz contains 9,000 frames; a simulated trajectory could instead contain 10,000 transitions. These are examples, not typical training-window sizes. Partitioning only the batch still leaves each assigned sequence on its device. For suitable architectures, context partitioning distributes token positions and their attention work, while some frameworks use sequence parallelism for selected activation storage alongside tensor parallelism. Neither term licenses ignoring temporal dependencies. We will make those meanings explicit below.

Data Parallelism and Gradient Synchronization

A simple baseline is fully replicated data parallelism: every device holds a full copy of the model, receives its assigned minibatch, computes gradients locally, and then synchronizes gradients before the optimizer step. Mathematically, if worker ii has local loss Li(θ)\mathcal{L}_i(\theta) computed from its own shard of the global batch, the goal is that every worker takes the same step using the global gradient

g=1N∑i=1N∇θLi(θ)g = \frac{1}{N} \sum_{i=1}^{N} \nabla_\theta \mathcal{L}_i(\theta)

where:

  • gg: the global gradient estimate, the average of the per-worker gradients
  • NN: the number of workers sharing the global batch
  • Li(θ)\mathcal{L}_i(\theta): the local loss computed by worker ii from its own shard of the global batch
  • ∇θLi(θ)\nabla_\theta \mathcal{L}_i(\theta): the gradient of worker ii's local loss with respect to the parameters θ\theta
  • θ\theta: the model parameters, identical on every worker before the optimizer step

The gradient exchange operation is called all-reduce: every rank contributes a local gradient and receives the reduced result. In ordinary replicated data parallelism, matching initial parameters, optimizer state, hyperparameters and update order let ranks apply the same synchronized update. Gradient synchronization does not synchronize arbitrary mutable model state or repair differing optimizer state. PyTorch DDP organizes synchronization into gradient buckets; it is not necessarily one collective per whole step.

For a loss that is the mean of independent per-example terms, equal-size shards justify the displayed unweighted mean. Unequal shard sizes require weights proportional to the number of contributing examples or valid tokens, with matching loss normalization. This is a mathematical equivalence, not a promise of bitwise equality: floating-point reduction order changes rounding, and batch-dependent operations such as local BatchNorm can change the computation itself. Learning-rate and batch-size choices still require validation rather than following automatically from distribution.

Two details decide efficiency. First, a billion fp32 gradient entries occupy 4 GB decimal locally; this is a payload size, not the bytes sent over the network. For the ideal ring algorithm below, each rank transmits 2(N−1)G/N2(N-1)G/N bytes for gradient payload GG. Other algorithms and compression change that count. Second, bucketing can overlap ready gradient reductions with remaining backward computation. Whether overlap hides much communication depends on bucket readiness, compute duration and network contention; a large model is not enough to guarantee it. Profile the particular encoder and dynamics workload.

Stochastic latent variables add a gradient-estimation question, not a requirement that data-parallel ranks sample identical latents. Independent random streams allow independent draws across ranks; after gradient reduction, parameters can still receive the same update. If the sampling distribution depends on parameters, differentiating a sampled value is not automatically an unbiased estimator. Valid pathwise or score-function estimators need their respective assumptions; straight-through estimators may be biased. Stochastic computation graphs make these dependencies explicit. Sharing seeds can correlate estimates and may increase or decrease variance depending on covariance; it does not generally buy lower variance. Replacing a sampled latent by its mean generally changes a nonlinear stochastic objective, although special cases such as affine costs can agree. Treat mean substitution as a modeling choice rather than a synchronization fix.

Tensor Parallelism, Pipeline Parallelism, and Their Costs

When layer or model storage exceeds one device's capacity, model-parallel schemes are possible remedies alongside state sharding or offload. Two useful families are:

  • Tensor parallelism (TP): partition operations inside a layer. For a column-partitioned matrix W=[W1,…,Wk]W=[W_1,\dots,W_k] in Y=XWY=XW, rank rr computes the output-feature shard Yr=XWrY_r=XW_r. Reconstructing Y=[Y1,…,Yk]Y=[Y_1,\dots,Y_k] requires concatenation, implemented by an all-gather when every rank needs the full output. A compatible next layer can consume the shards without gathering them first. In a row-partitioned multiply, each rank instead computes a contribution to the same output entries, and the contributions must be summed, for example with an all-reduce or a reduce-scatter. Megatron's linear-layer implementation makes this distinction explicit. Sharding reduces local weight storage, but layer-internal communication can become latency-critical. Its collective count and overlap depend on shard layouts, scheduling and kernels, not simply the number of layers. Fast links are valuable where communication is exposed; profiling determines the useful rank count.
  • Pipeline parallelism (PP): consecutive layer blocks are assigned to stages, and independent microbatches flow through them. One stage can process a later microbatch while the next processes an earlier one. Filling and draining leave idle stage slots. For an ideal balanced forward pipeline with PP stages, MM microbatches, equal-duration stage work, no communication overhead and no overlap between separate batches, each stage does MM active slots within M+P−1M+P-1 elapsed slots. Its idle fraction is
bubble fraction=P−1M+P−1\text{bubble fraction} = \frac{P-1}{M+P-1}

where:

  • PP: the number of pipeline stages (each stage is a contiguous block of layers on its own device)
  • MM: the number of microbatches into which each global batch is split
  • P−1P - 1: idle slots for an individual stage over the complete forward schedule; the split between filling and draining depends on the stage's position

At fixed PP, increasing MM makes this ideal idle fraction approach zero. It does not eliminate real communication, imbalance or memory costs. Training adds backward work and may use a different schedule, so apply the formula only to a schedule for which its slot accounting holds. The conveyor-belt intuition is useful: more independent items amortize startup and teardown, provided stations can process them at compatible rates.

For a large video encoder, tensor parallelism may address a layer-size limit; deep stacks can provide pipeline partitions. Rollout training still needs a workload trace before choosing either. A temporal transition is not necessarily an optimizer update or synchronization boundary: a learner can form one loss from an entire batch of imagined trajectories. Independent environments or particles can supply microbatches, but causal dependencies inside each trajectory remain. When utilization is poor, measure stage imbalance, communication, memory and schedule bubbles. Larger independent batches are one possible remedy, not a guaranteed cure or a requirement of all imagination-based training.

Sequence Parallelism for Long-Horizon Rollouts

Long-context training can exceed activation memory even when parameters fit. Partitioning known token positions distributes some of that storage and computation. The relevant length is the actual training context, not the total duration of the source video or a hypothetical imagination horizon. Tokens, frames and environment steps are different units, and there is no family-wide rule that world models use longer contexts than language models. The key question is which operations can be distributed without changing the computation.

The semantics differ by architecture. For a recurrent state-space model with hidden state hth_t, the recurrence

ht=f(ht−1,xt)h_t = f(h_{t-1}, x_t)

where:

  • hth_t: the hidden state at time tt
  • ht−1h_{t-1}: the hidden state at the previous time step
  • xtx_t: the input at time tt
  • ff: the deterministic state-transition function (the recurrent core)

has a causal dependency: a shard needs the preceding shard's final state. For an arbitrary nonlinear recurrent core, computing each chunk from zero and combining its final states cannot reconstruct the true trajectory. Parallel scans require a compact associative representation of transition composition. For example, input-conditioned affine maps ht=Atht−1+bth_t=A_t h_{t-1}+b_t compose as (A2A1,A2b1+b2)(A_2A_1,A_2b_1+b_2); whether that representation is efficient depends on matrix structure and whether coefficients can be computed without the previous state. This does not establish an efficient scan for a nonlinear Dreamer-style recurrence. Such a model can batch independent episodes or particles while preserving time order within each trajectory; time partitioning needs actual boundary states, recomputation or an explicitly approximate training scheme.

For a transformer world model, context parallelism partitions token positions. In Ring Attention, queries stay on their owning device while key/value blocks circulate. Blockwise attention combines partial results using online softmax statistics rather than materializing the full attention matrix. Causal masks prevent attending to future positions. Dense attention still has quadratic total arithmetic, while token storage is distributed and temporary score memory is bounded by the chosen block size. Long-context attention over already available inputs does not remove autoregressive dependencies when generating new trajectory tokens. NVIDIA also distinguishes sequence parallelism from context parallelism: the former in Megatron shards selected activations alongside tensor parallelism, while the latter partitions the input context and attention work.

For an ordinary sequential diffusion or flow sampler, each solver step consumes the previous step's state. Assigning successive dependent steps to different devices therefore requires transfers at their boundaries, not necessarily a global collective. One exact baseline keeps solver steps ordered while distributing within-step spatial-temporal token work. Independent samples can also be batched. Alternative parallel or approximate solvers require their own convergence and quality analysis; the sequential dependency alone does not prove that every other distributed schedule is inefficient.

Let's compare the communication patterns to build intuition:

  • Replicated data parallel: gradient-bucket collectives per synchronized backward pass, with payload proportional to gradient entries; overlap depends on the workload.
  • Tensor parallel: layer-internal gathers, reductions or scatters, determined by partition layout. Count each collective's payload and rounds over the actual forward/backward schedule; some dependencies expose latency while others can overlap work.
  • Pipeline parallel: point-to-point activation transfers at stage boundaries for each microbatch, with gradient transfers in backward. For PP stages and MM microbatches, the basic forward schedule has M(P−1)M(P-1) boundary transfers; their byte counts depend on the boundary tensors. More microbatches create overlap opportunities but do not guarantee hidden traffic.
  • Ring context parallel attention: an ideal unpruned forward sweep over NN ranks can visit all KV blocks with N−1N-1 sends per rank, without a final return exchange. Each block has size proportional to SdKV/NS d_{\mathrm{KV}}/N, so this schedule transmits per-rank KV entries proportional to (N−1)SdKV/N(N-1)S d_{\mathrm{KV}}/N, linear in total tokens SS, not quadratic. Here dKVd_{\mathrm{KV}} is the combined key/value width. Count the actual implementation's exchanges: the paper's Appendix A forward loop also exchanges after the last computation, making NN sends. Backward communication and causal pruning need separate accounting.

These patterns trade bytes, synchronization frequency and dependency placement. There is no universal ranking in which data parallelism always moves the most bytes or tensor parallelism always moves less. Model width, context length, microbatch size and rank count determine the comparison. Ring attention's quadratic arithmetic and linear per-rank forward KV traffic are different quantities; chunking bounds working memory rather than converting quadratic network traffic into linear traffic.

Large-model systems can combine several parallelism axes. A combination of data, tensor and pipeline parallelism is often termed 3D parallelism; an additional context or other partition introduces another axis, but these names do not define one mandatory world-model recipe. The mixture depends on interconnect, tensor shapes, sequence length and the measured schedule. A useful heuristic is to put the most exposed communication on the fastest available links. Data-parallel reductions can span slower links when enough backward work or gradient accumulation amortizes their cost; short synchronized updates can be latency-sensitive too.

An Illustrative Model of Gradient Sync Cost

To build intuition about gradient synchronization, consider an idealized ring all-reduce. A reduce-scatter phase and an all-gather phase each take N−1N-1 rounds, with G/NG/N bytes sent per rank in each round. This gives the following communication-only model for N>1N>1:

T(N)=2(N−1)α+2(N−1)NβGT(N) = 2(N-1)\alpha + \frac{2(N-1)}{N}\beta G

where:

  • T(N)T(N): the time to complete the all-reduce collective with NN workers
  • α\alpha: latency per communication round in seconds
  • β\beta: the per-byte transfer cost of the interconnect
  • GG: gradient bytes held by each rank before the collective
  • NN: worker count; a one-worker job has no synchronization, so T(1)=0T(1)=0

This derivation counts transmitted bytes and rounds, not reduction arithmetic, contention, or compute/communication overlap. Ring algorithms are analyzed by Patarasuk and Yuan (2009). The values below are hypothetical, not measurements of a GPU cluster. Holding global work fixed, we assume computation takes 3/N3/N seconds per worker and add the unoverlapped synchronization time. The byte-transfer term approaches 2βG2\beta G as workers increase, while the round-latency term grows.

In[3]:
Code
import time

import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F

# Hypothetical ring communication model; no cluster is benchmarked.
ALPHA = 40e-6  # 40 microseconds
BETA = 0.6e-9  # 0.6 ns per byte
PARAMS = 100e6  # 100M-parameter model
BYTES_PER_PARAM = 4  # fp32 gradients
GRAD_BYTES = PARAMS * BYTES_PER_PARAM  # 400 MB gradient

worker_counts = np.array([1, 2, 4, 8, 16, 32, 64, 128, 256])
sync_times = (
    2 * (worker_counts - 1) * ALPHA
    + 2 * (worker_counts - 1) / worker_counts * BETA * GRAD_BYTES
)
compute_per_worker = (
    3.0 / worker_counts
)  # total 3 seconds of work, split evenly

eff_ratio = compute_per_worker / (compute_per_worker + sync_times)
Out[4]:
Console
workers |  sync ms |  compute s |  efficiency
      1 |    0.000 |     3.0000 |     100.00%
      2 |  240.080 |     1.5000 |      86.20%
      4 |  360.240 |     0.7500 |      67.55%
      8 |  420.560 |     0.3750 |      47.14%
     16 |  451.200 |     0.1875 |      29.36%
     32 |  467.480 |     0.0938 |      16.70%
     64 |  477.540 |     0.0469 |       8.94%
    128 |  486.410 |     0.0234 |       4.60%
    256 |  498.525 |     0.0117 |       2.30%

The table starts at 100% compute efficiency for one worker, because no collective is needed. Efficiency then decreases in this strong-scaling model because per-worker computation shrinks, but the communicated gradient does not shrink to zero. Efficiency is the fraction of the modeled step spent computing, not a measured utilization or a direct speedup. Actual hardware may use other collective algorithms and overlap; those must be profiled rather than inferred from this toy.

The toy identifies a conditional bottleneck: when exposed synchronization takes longer than local computation, it occupies most of the modeled step. A rollout's individual temporal transitions need not each trigger a collective. Larger independent batches, gradient accumulation or better overlap can improve the computation-to-communication ratio, subject to memory and learning constraints. Check actual optimizer boundaries and traces rather than infer synchronization frequency from rollout length alone.

The modeled efficiency ratio shows the communication penalty as fixed global work is divided among more workers.

Out[5]:
Visualization
Modeled compute efficiency decreases as more workers share fixed global work.
Modeled compute efficiency with fixed global work and unoverlapped ring all-reduce. Compute time shrinks as workers increase, while synchronization retains a byte-transfer cost and gains more communication rounds. These are hypothetical costs, not measured cluster results.

Data, Checkpoint, and Experiment Infrastructure

Distributed training supplies compute; data and artifact infrastructure supply the intended inputs and evidence of what happened. This section connects data loading, checkpointing and experiment tracking. Exact continuation requires the causally relevant training state to be captured or reconstructible. A stochastic loader can be replayable if its random streams, ordering and pending work are recorded. An incomplete checkpoint may still support some comparisons, but cannot justify exact continuation. Code identity can come from a commit, a content hash or an archived executable snapshot; losing every usable copy prevents faithful reconstruction. These records also support audits, without establishing legal compliance by themselves.

Data Loading for World-Model Trajectories

For action-conditioned reinforcement-learning setups, training trajectories can contain temporally aligned (observation,action,reward)(\text{observation}, \text{action}, \text{reward}) sequences. Two properties distinguish such trajectory data from typical image classification pipelines:

  • Temporal association. Dynamics learning needs evidence linking states, actions and subsequent outcomes. Sampling isolated frames while discarding their episode and time associations removes those observed transition pairs, even though individual images can still support representation learning or prior knowledge. A sampling unit can be a contiguous window or an episode; its length is a workload choice, not a fixed rule.
  • Heterogeneous modalities. A sample may combine RGB, depth, proprioception, torques, actions, language and metadata. Their decoding and transforms can differ. Whether implemented as separate paths or one dispatching loader, the requirement is to preserve temporal alignment and sample identity across the streams.

A loader can be organized as an assembly line: retrieve and decode data, sample contiguous windows, transform aligned modalities, collate tensors and masks, then prefetch the next batches. Record augmentation randomness when replay matters. Choose prefetch depth from service times, memory and recovery requirements rather than always buffering the same number of batches. Any stage can become the bottleneck; measure where the consumer waits.

The main failure patterns here:

  • Shard hotspots. Assigning entire unequal-work shards to workers can create imbalance; unequal file sizes alone do not prove it if workers process equal-cost windows. Balance assignments using measured decode and sample cost, and distinguish load-balancing weights from statistical sampling weights. Exposed synchronization waits for required ranks, so an overloaded rank can limit throughput. Overlap and other serial work determine the magnitude.
  • Repeated retrieval overhead. Many small requests and repeated decoding can cost more than fetching and reusing a shard. Whole-shard reads can amortize overhead when many windows reuse the bytes; indexed range reads can be preferable when little of a large shard is needed. Benchmark request concurrency, locality, transferred bytes and decode reuse rather than assume one access pattern is always cheaper.
  • Decode bottlenecks. When the decoder cannot supply batches at the model's consumption rate, it can starve the accelerator. Pre-decoding, supported hardware decoders or parallel CPU workers are possible remedies, with storage and format tradeoffs. A large slow model may instead be the bottleneck, so buffering and parallelism should follow measurements.

Determinism needs more than a dataset seed. Record the dataset manifest and version, sampler algorithm, epoch and cursor, augmentation RNG states, rank count and work assignment, plus any streaming offsets or prefetched items that affect the next batch. A different shuffle order is a different realized optimization sequence, not necessarily a different underlying data distribution. Framework kernels, releases and hardware can also affect numerical reproducibility, as PyTorch's reproducibility notes explain. The restricted example below tests a fixed CPU setting; it does not promise replay across arbitrary worker counts or platforms.

Checkpointing

A checkpoint serializes the state needed to continue training: parameters, optimizer moments, relevant schedules, random streams, data position and any auxiliary network state the algorithm uses. Under a fixed replayable configuration, restoring complete state can reproduce subsequent updates. Missing state can repeat data or change updates without an immediate error. That motivates a continuation test near the save/load code rather than relying on a plausible-looking loss curve.

Three concepts are essential:

  • Atomic publication and durability. Publish a checkpoint as current only after every required snapshot file is complete and validated. A partial file may fail to load, or may pass inadequate checks. On a supporting local filesystem, a same-filesystem temporary-file replacement can provide atomic visibility; it does not by itself prove persistence after power loss. Durable recovery additionally depends on file and directory synchronization or the storage service's committed-write guarantees. Distributed shards need a protocol that publishes one complete snapshot identity, such as an atomically committed manifest, not a promise that every storage system supports rename.
  • Consistency across ranks. Snapshot one agreed logical optimizer-step boundary, including all sharded states. A backward all-reduce does not establish that every subsequent optimizer update or asynchronous write has completed. Complete the relevant updates, coordinate ranks and capture an immutable consistent snapshot before committing the checkpoint manifest. Replicated parameters and sharded state require different save protocols.
  • Tiering and pruning. One retention policy saves every NN steps, retains the last KK, and promotes selected snapshots to longer-term storage. Choose retention from recovery, reproducibility and audit requirements; keeping every snapshot can be justified. Provenance comes from an explicit deployed checkpoint identity and traceable records, not from deleting alternatives. Storage tier and retained recovery points affect retrieval cost and lost work.

Two gotchas specific to world models:

  • Data loader state. Preserve enough state to reconstruct the next data selection, not just the model tensors. Replaying or skipping windows changes the realized update sequence, although it need not introduce statistical bias in every sampling scheme. Some pipelines use a deterministic global sampler keyed on step so the loader state is reconstructible from a step number and unchanged dataset and sampler configuration. In a streaming or prefetched loader, reconstructing the next selection may also require pending work and random-stream state. Model and data state together determine continuation.
  • Auxiliary learning state. Save whichever target networks, slow critics, EMA copies, normalization statistics and update counters the specific algorithm actually uses. Dreamer, TD-MPC and MuZero variants do not all maintain the same EMA copy of the entire model. A missing auxiliary tensor or schedule position can change learning after resumption even if the primary parameters load correctly. Audit the implemented algorithm rather than infer checkpoint contents from its family name.

Checkpoint size depends on the actual saved tensors. As an illustration, suppose a checkpoint stores fp16 weights, a separate fp32 master copy and two fp32 Adam moments. That is 2+4+82 + 4 + 8 bytes per parameter, excluding auxiliary tensors and serialization overhead. The separate master-weight approach is described in Mixed Precision Training. It is not universal: ordinary PyTorch autocast can retain fp32 parameters without storing an additional fp16 copy. The tutorial below uses fp32 parameters and AdamW. Activations and gradients normally are not saved in an optimizer-step checkpoint. For the stated illustrative layout:

Mcore=Nparams⋅(bweights+bmaster+bmoments)(definition)=109⋅(2+4+8) bytes(substitute; note: 2+4+8=14)=14×109 bytes(simplify)≈13.04 GiB(divide bytes by 230)\begin{aligned} M_{\text{core}} &= N_{\text{params}} \cdot (b_{\text{weights}} + b_{\text{master}} + b_{\text{moments}}) && \text{(definition)} \\ &= 10^9 \cdot (2 + 4 + 8) \text{ bytes} && \text{(substitute; note: } 2 + 4 + 8 = 14 \text{)} \\ &= 14 \times 10^9 \text{ bytes} && \text{(simplify)} \\ &\approx 13.04 \text{ GiB} && \text{(divide bytes by } 2^{30}\text{)} \end{aligned}

where:

  • McoreM_{\text{core}}: total core training state in bytes (weights, master copy, and optimizer moments)
  • NparamsN_{\text{params}}: the parameter count, here 10910^9
  • bweights=2b_{\text{weights}} = 2: bytes per parameter for fp16 weights
  • bmaster=4b_{\text{master}} = 4: bytes per parameter for the fp32 master copy
  • bmoments=8b_{\text{moments}} = 8: bytes per parameter for the two fp32 Adam moments

The stated tensors occupy 14 GB decimal, about 13.04 GiB, before auxiliary buffers or file overhead; this conversion does not round buffers into the total. Repeated writes create an I/O budget. Options include asynchronous checkpointing with an immutable snapshot and a durable completion marker, sharded checkpoints with a committed manifest, and delta checkpointing when actual encoded deltas reduce bytes. Parameter values changing slowly does not by itself make dense checkpoints smaller.

Experiment Infrastructure

Experiment infrastructure is the layer that makes a training run an experiment rather than merely a computation. Minimally, it must record:

  • Configuration: model architecture, hyperparameters, seeds, code version (git commit or content hash), container image digest.
  • Data: dataset version, shard manifest, preprocessing script version, licensing/consent tags for any human data.
  • Metrics: training loss curves, wall-clock per step, throughput (tokens or frames per second), gradient norm, effective learning rate, memory headroom, and validation metrics captured per checkpoint.
  • Artifacts: checkpoints, evaluation rollouts, any visualization dumps (a critical debugging aid for latent world models, see Interpretability and World-Model Debugging).
  • Environment: accelerator type, interconnect, driver/CUDA versions, and whether mixed precision, sequence parallelism, or activation checkpointing were enabled.

These records make comparisons interpretable. A difference between two evaluation scores does not establish that one intended change caused it. Record potential confounders, use controlled comparisons and replication where appropriate, and report uncertainty. Different realized sample orders can contribute run-to-run variability even when the dataset is unchanged. Missing provenance may require rerunning an experiment, but logging is evidence for a comparison, not proof that every other factor was identical.

A Minimal Checkpoint and Data Pipeline

The following code demonstrates a scheduler-free, single-device checkpoint loop that captures the essential state, plus a sharded, deterministic data sampler. The goal is clarity, not performance, and it runs comfortably on CPU. Read it with an eye to which pieces of state are captured and why each one is necessary for a resumed run to match an uninterrupted one.

In[6]:
Code
class TinyLatentDynamics(nn.Module):
    """Minimal latent world model: predicts next latent from (z, a)."""

    def __init__(self, latent_dim=16, action_dim=2, hidden=64):
        super().__init__()
        self.latent_dim = latent_dim
        self.action_dim = action_dim
        self.net = nn.Sequential(
            nn.Linear(latent_dim + action_dim, hidden),
            nn.ReLU(),
            nn.Linear(hidden, hidden),
            nn.ReLU(),
            nn.Linear(hidden, latent_dim),
        )

    def forward(self, z, a):
        return self.net(torch.cat([z, a], dim=-1))


class DeterministicWindowSampler:
    """Each worker gets a deterministic, non-overlapping shard of windows."""

    def __init__(
        self,
        n_episodes,
        episode_len,
        window_len,
        stride,
        rank=0,
        world_size=1,
        seed=1234,
    ):
        for name, value in (
            ("n_episodes", n_episodes),
            ("episode_len", episode_len),
            ("window_len", window_len),
            ("stride", stride),
            ("world_size", world_size),
        ):
            if (
                isinstance(value, bool)
                or not isinstance(value, (int, np.integer))
                or value <= 0
            ):
                raise ValueError(f"{name} must be a positive integer")
        if (
            isinstance(rank, bool)
            or not isinstance(rank, (int, np.integer))
            or not 0 <= rank < world_size
        ):
            raise ValueError("rank must be an integer in [0, world_size)")
        if window_len > episode_len:
            raise ValueError("window_len cannot exceed episode_len")
        windows = []
        for ep in range(n_episodes):
            for start in range(0, episode_len - window_len + 1, stride):
                windows.append((ep, start))
        rng = np.random.default_rng(seed)
        rng.shuffle(windows)
        self.windows = windows[rank::world_size]

    def __len__(self):
        return len(self.windows)

    def __getitem__(self, i):
        return self.windows[i]

We can now define checkpoint state for this restricted CPU example. The step count identifies the next window only while the dataset, sampler seed, rank count and sampler algorithm remain unchanged. It is not sufficient for an arbitrary streaming loader. The helper assumes one writer, no concurrent training update, and a local filesystem supporting same-directory atomic replacement. Its fixed temporary name is not a concurrent-writer protocol, and it calls neither file nor directory fsync, so it is not a crash-durability test. Distributed storage requires its own consistency, commit and durability protocol. Load only a trusted checkpoint written by this example, because weights_only=False enables pickle deserialization.

In[7]:
Code
def save_checkpoint(path, model, optimizer, step, sampler_seed):
    checkpoint_path = Path(path)
    temporary_path = checkpoint_path.with_suffix(
        checkpoint_path.suffix + ".tmp"
    )
    torch.save(
        {
            "step": step,
            "weights": model.state_dict(),
            "optimizer": optimizer.state_dict(),
            "rng_torch": torch.get_rng_state(),
            "rng_numpy": np.random.get_state(),
            "sampler_seed": sampler_seed,
        },
        temporary_path,
    )
    temporary_path.replace(checkpoint_path)


def load_checkpoint(path, model, optimizer):
    state = torch.load(path, weights_only=False)
    model.load_state_dict(state["weights"])
    optimizer.load_state_dict(state["optimizer"])
    torch.set_rng_state(state["rng_torch"])
    np.random.set_state(state["rng_numpy"])
    return state["step"], state["sampler_seed"]

Let's generate a controlled synthetic state sequence, with 16 state dimensions and two action dimensions. Each episode follows st+1=0.8st+0.2atB+ϵts_{t+1}=0.8s_t+0.2a_tB+\epsilon_t, where BB has shape (2,16)(2,16) and independent process noise has standard deviation 0.02. The observations here equal the synthetic states; no encoder or learned state representation is being tested. The model fits one-step transitions, not a validated real-world simulator.

In[8]:
Code
model = TinyLatentDynamics()
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3)

# Controlled synthetic dynamics, not independently sampled next-state targets.
T_gen, n_ep = 32, 64
sim_rng = np.random.default_rng(0)
actions = sim_rng.normal(size=(n_ep, T_gen, 2)).astype("float32")
action_map = sim_rng.normal(scale=0.5, size=(2, 16)).astype("float32")
latent_obs = np.zeros((n_ep, T_gen + 1, 16), dtype="float32")
latent_obs[:, 0] = sim_rng.normal(size=(n_ep, 16))
process_noise = sim_rng.normal(scale=0.02, size=(n_ep, T_gen, 16))
for t in range(T_gen):
    latent_obs[:, t + 1] = (
        0.8 * latent_obs[:, t]
        + 0.2 * actions[:, t] @ action_map
        + process_noise[:, t]
    )

sampler = DeterministicWindowSampler(
    n_episodes=n_ep, episode_len=T_gen + 1, window_len=8, stride=1, seed=7
)


def train_window(model, optimizer, sampler, window_idx):
    ep, start = sampler[window_idx]
    z = torch.from_numpy(latent_obs[ep, start : start + 7])
    a = torch.from_numpy(actions[ep, start : start + 7])
    z_next = torch.from_numpy(latent_obs[ep, start + 1 : start + 8])

    pred = model(z, a)
    loss = F.mse_loss(pred, z_next)
    optimizer.zero_grad()
    loss.backward()
    optimizer.step()
    return float(loss.detach())


checkpoint_step = 512
losses = [
    train_window(model, optimizer, sampler, i) for i in range(checkpoint_step)
]
save_checkpoint(
    "tiny_wm.ckpt.pt", model, optimizer, checkpoint_step, sampler_seed=7
)

Now compare the next 16 uninterrupted updates with a fresh model and optimizer loaded from that checkpoint. Both paths use the same CPU runtime and dataset. Equal losses alone are insufficient, so we also compare all resulting model parameters. We additionally sample and compare the next Torch CPU and legacy NumPy random draws; the deterministic training continuation alone would not exercise those saved random streams.

In[9]:
Code
continuation_steps = 16
expected_torch_draw = torch.rand(4)
expected_numpy_draw = np.random.standard_normal(4)
expected_losses = [
    train_window(model, optimizer, sampler, i)
    for i in range(checkpoint_step, checkpoint_step + continuation_steps)
]
expected_weights = {
    name: value.clone() for name, value in model.state_dict().items()
}
resumed_model = TinyLatentDynamics()
resumed_optimizer = torch.optim.AdamW(resumed_model.parameters(), lr=1e-3)
resumed_step, resumed_seed = load_checkpoint(
    "tiny_wm.ckpt.pt", resumed_model, resumed_optimizer
)
assert torch.equal(expected_torch_draw, torch.rand(4))
assert np.array_equal(expected_numpy_draw, np.random.standard_normal(4))
resumed_sampler = DeterministicWindowSampler(
    n_ep, T_gen + 1, 8, 1, seed=resumed_seed
)
resumed_losses = [
    train_window(resumed_model, resumed_optimizer, resumed_sampler, i)
    for i in range(resumed_step, resumed_step + continuation_steps)
]
assert np.array_equal(expected_losses, resumed_losses)
assert all(
    torch.equal(expected_weights[name], value)
    for name, value in resumed_model.state_dict().items()
)
model = resumed_model
optimizer = resumed_optimizer
Out[10]:
Console
Available windows: 1664; checkpoint after 512 updates
Final loss: 0.0038
Mean loss over first 10 steps: 0.0889
Mean loss over last 10 steps: 0.0552
Exact loss and weight match for 16 CPU continuation updates

Compare the printed first-ten and last-ten means rather than assuming every noisy window loss decreases. Matching future losses and weights exercises the restored optimizer and sampler on the next 16 updates; it is not a direct equality assertion over every optimizer field. The next-draw assertions cover the two random streams saved here, not an independent default_rng generator or a device RNG. This is not a distributed-recovery proof: topology changes, prefetch state, schedules, scalers, interrupted writes and different software or hardware require additional tests. PyTorch's reproducibility guidance limits cross-platform and cross-release reproducibility.

The pre-checkpoint loss curve shows optimization on a controlled synthetic process, including variation between sampled windows. The separate continuation assertions, not the shape of this plot, test resumption.

Out[11]:
Visualization
One-step training loss varies between synthetic trajectory windows.
One-step MSE over training windows from controlled synthetic dynamics before checkpointing. Window-to-window variation remains; the separate continuation test checks exact restored losses and parameters in the fixed CPU runtime.

Rollout Serving, Caching, and Batching

One deployment option is a service performing autoregressive rollouts. A controller may submit a request every 20 milliseconds, while an offline evaluator or imagination-based learner may evaluate thousands of independent candidates. These are illustrative workloads. They can use separate stacks; if they share capacity, serving must account for both interactive deadlines and batch throughput rather than assume one policy suits both.

For autoregressive models, two properties drive serving design:

  • Causal state. A recurrent transition can consume a fixed-size sufficient state and the current action. A full-context causal transformer can instead retain a growing token history or KV cache. State must have an owner and identity, but that owner can be the client, server or local controller; model statefulness does not require server-held sessions.
  • Decode granularity. Some models emit one token or a small block per call; others predict a joint trajectory or video. Small operations can expose dispatch overhead, whereas a large encoder or denoiser can dominate computation. Measure the actual boundary and choose batching against throughput and deadlines.

Rollout Serving Architectures

Two possible API designs are:

Stateless request/response (POST /rollout). A request supplies enough state or observation context, actions and inference configuration to produce its result. Any compatible worker can handle that self-contained request, which simplifies routing but does not remove model-placement, workload or deadline constraints. It can serve offline or interactive workloads. A caller holding a sufficient recurrent state can send that compact state and one action rather than full history. A transformer API that requires retransmitting a long context can have substantial transfer and recomputation costs.

Stateful session (POST /session/{id}/step). The server owns the current state and any cache; each ordered call advances that session. This can avoid repeatedly transferring context for interactive clients. Define session lifetimes, capacity, model-version binding, step ordering, retries and routing to the owner or a correct state-transfer mechanism. An eviction must return an explicit expiry or use a valid recovery protocol, not silently continue from a default state.

A system can offer both designs and route requests by declared workload or latency class. That is an architecture option, not a requirement or a measured prevalence claim. Triton's sequence batcher illustrates one stateful implementation: requests in a sequence are routed to the same model instance.

Caching

Caching in a world-model server comes in three categories, each solving a different problem.

Key-value (KV) cache. For causal transformer decoding with one new token per call and fixed layer count and width, recomputing a dense-attention prefix of length tt costs O(t2)O(t^2) attention arithmetic per call. Reusing past keys and values leaves the new query attending to tt cached positions, costing O(t)O(t) attention arithmetic. Over TT generated tokens, those attention totals are O(T3)O(T^3) versus O(T2)O(T^2), not quadratic versus linear. Cached decoding remains increasingly expensive with unbounded context. A fixed window of WW positions gives O(TW)O(TW) attention arithmetic, but changes available context unless the model was defined for that window. Token count is not environment-step count when each observation contains several tokens; projections and other layers add their own costs. KV caching avoids recomputing old keys and values, not the new query's attention to retained history.

For one session with ordinary multi-head attention, where the number of KV heads equals the query-head count, cache tensor storage grows linearly with cached token count. Ignoring allocator overhead, its size is

MKV=2⋅L⋅H⋅dh⋅S⋅2 bytesM_{\text{KV}} = 2 \cdot L \cdot H \cdot d_h \cdot S \cdot 2 \text{ bytes}

where:

  • LL: the number of transformer layers
  • HH: the number of KV heads per layer, equal to query heads in this MHA example; grouped-query or multi-query attention has fewer KV heads
  • dhd_h: the head dimension
  • SS: cached token positions, not necessarily environment time steps when a frame has multiple tokens
  • the leading 22: one key tensor and one value tensor per layer/head
  • the trailing 22: bytes per fp16 scalar

With L=32L=32, H=16H=16, dh=64d_h=64 and S=8192S=8192 token positions, the cache tensors occupy 1,073,741,8241,073,741,824 bytes: exactly 1 GiB or about 1.074 GB decimal per session. Multiply this tensor count by concurrent sessions before adding allocator overhead. Sliding-window attention and cache eviction may reduce stored history but require explicit semantics for lost context; they are not automatically equivalent to full-context inference. Cache layout and length-aware scheduling affect actual allocation behavior.

Result caching. Memoization helps when requests repeat a complete, deterministic computation, but the benefit depends on hit rate and lookup cost rather than a universal speedup range. Include model/checkpoint version, preprocessing and inference configuration, initial state or observation history, exact action prefix, horizon and numeric execution policy in the key. For seeded sampling, include the RNG algorithm and explicit stream state or a stable keyed counter; a seed alone does not identify a draw after arbitrary consumption. Reusing one cached stochastic result is not equivalent to requesting independent samples. Either disable memoization for fresh-sample requests or define a reproducible sample identity in the API and key. Invalidate on model or configuration changes, and never share session state across independent particles merely because their visible inputs match.

Retrieved-context caching. As an illustrative design, a map-conditioned model could retrieve a region's features before predicting local dynamics. Reuse is valid only when region content, query, preprocessing and model-dependent representation remain compatible. A region identifier alone is not enough after the map or representation changes. Whether caching helps depends on hit rate, retrieval cost, lookup overhead and invalidation; this tutorial does not benchmark a hierarchical retrieval system.

Batching

Batching groups independent requests into a forward pass to amortize overhead. Iteration-level or continuous batching can admit new sequences as others finish, as in Orca. Rollout sessions can be at different indices with different retained lengths, so ragged serving needs correct cache ownership, offsets and eligible-token masks. In ordinary attention with a separate batch row per session, the batch dimension isolates rows; omitting a padding mask does not by itself read another row. Packed sequences within one token axis additionally need boundary masking, and causal masking excludes future tokens. Masks cannot repair a cache routed to the wrong session.

The main tradeoffs of batching:

  • Larger batches can amortize fixed dispatch overhead, but extra queue waiting, work and memory can increase response time. Their measured latency need not grow monotonically with size. A batching wait that exceeds a controller's remaining slack misses its deadline.
  • Smaller batches can be appropriate for low demand or tight deadlines. Their utilization depends on operation sizes and hardware, not batch size alone.
  • Length-aware grouping reduces padding in a representation padded to its longest member. It is not chunked prefill: that technique splits the processing of one input context into smaller pieces. Packed or paged layouts have different waste and allocation behavior.

A scheduler may deliberately wait for more requests up to a bounded deadline, or dispatch immediately. Choose that bound from arrivals, service times and available slack; there is no general few-microsecond setting. The packet-buffer analogy is useful for thinking about delay versus utilization, but does not establish correctness. Triton exposes an optional maximum queue delay for dynamic batching; a deployment still needs end-to-end measurements.

A Tiny Stateful Rollout Server

The next block builds a single-process, single-threaded in-memory server with state per session and FIFO eviction by creation order. It requires positive integer capacities and valid tensor shapes, copies externally supplied state and returned results, and raises KeyError for missing or evicted IDs. It does not implement LRU, timeouts, batching, locks, authenticated requests or idempotent retry/step ordering. History is bookkeeping only, not an attention KV cache; the fixed model consumes current state and action.

In[12]:
Code
class LatentHistory:
    """Fixed-capacity rolling history of past latents (a stand-in for a real KV cache)."""

    def __init__(self, latent_dim, max_len):
        for name, value in (("latent_dim", latent_dim), ("max_len", max_len)):
            if (
                isinstance(value, bool)
                or not isinstance(value, (int, np.integer))
                or value <= 0
            ):
                raise ValueError(f"{name} must be a positive integer")
        self.max_len = max_len
        self.latent_dim = latent_dim
        self.buffer = None  # (max_len, latent_dim)

    def append(self, z):
        if z.shape != (self.latent_dim,):
            raise ValueError("history entry must have shape (latent_dim,)")
        z = z.detach().clone()
        if self.buffer is None:
            self.buffer = z.unsqueeze(0)
        else:
            self.buffer = torch.cat([self.buffer, z.unsqueeze(0)], dim=0)
        if self.buffer.shape[0] > self.max_len:
            self.buffer = self.buffer[-self.max_len :].clone()


class RolloutSession:
    def __init__(self, session_id, z0, history_len=64):
        if z0.ndim != 2 or z0.shape[0] != 1:
            raise ValueError("initial state must have shape (1, latent_dim)")
        self.id = session_id
        self.z = z0.detach().clone()
        self.history = LatentHistory(z0.shape[-1], history_len)

    def step(self, model, action):
        if action.shape != (model.action_dim,):
            raise ValueError("action must have shape (action_dim,)")
        with torch.no_grad():
            z_next = model(self.z, action.unsqueeze(0))
        self.z = z_next
        self.history.append(self.z.squeeze(0))
        return self.z.detach().clone()


class RolloutServer:
    def __init__(self, model, max_sessions=16, history_len=64):
        for name, value in (
            ("max_sessions", max_sessions),
            ("history_len", history_len),
        ):
            if (
                isinstance(value, bool)
                or not isinstance(value, (int, np.integer))
                or value <= 0
            ):
                raise ValueError(f"{name} must be a positive integer")
        self.model = model
        self.max_sessions = max_sessions
        self.history_len = history_len
        self.sessions = {}
        self.counter = 0

    def open(self, z0):
        if z0.shape != (1, self.model.latent_dim):
            raise ValueError(
                "initial state must have shape (1, model.latent_dim)"
            )
        if len(self.sessions) >= self.max_sessions:
            oldest = min(self.sessions.values(), key=lambda s: s.id)
            del self.sessions[oldest.id]
        self.counter += 1
        sid = self.counter
        self.sessions[sid] = RolloutSession(sid, z0, self.history_len)
        return sid

    def step(self, sid, action):
        return self.sessions[sid].step(self.model, action)

We can now time a serial sweep over 16 independent sessions. Each session performs a separate forward call; this is not a vectorized batch. The later experiment measures vectorized batches separately.

In[13]:
Code
model.eval()
server = RolloutServer(model, max_sessions=64, history_len=128)

z0 = torch.zeros(1, model.latent_dim)
num_sessions = 16
session_ids = [server.open(z0.clone()) for _ in range(num_sessions)]

# Warm up
for _ in range(5):
    for sid in session_ids:
        server.step(sid, torch.zeros(model.action_dim))

step_times = []
for _ in range(200):
    t0 = time.perf_counter()
    for sid in session_ids:
        server.step(sid, torch.zeros(model.action_dim))
    step_times.append(time.perf_counter() - t0)

step_times = np.asarray(step_times)
print(f"Serial sweep of {num_sessions} sessions, {len(step_times)} repetitions")
print(f"Median sweep time: {np.median(step_times) * 1e3:.3f} ms")
print(
    f"Amortized time per forward: {np.median(step_times) / num_sessions * 1e3:.4f} ms"
)

Dividing the sweep time by session count reports amortized serial work, not the latency experienced by each caller and not a batching speedup. Session position in the serial sweep, admission queues and other work affect actual completion time.

The histogram records local serial-sweep timing variability. Its median is not a deadline guarantee: production decisions require end-to-end tail latency under representative concurrent load.

Out[14]:
Visualization
Histogram of local CPU time for serial sweeps over 16 sessions.
Local CPU timings for serial sweeps over 16 sessions. Each session runs a separate forward call. The histogram describes this microbenchmark, not vectorized batching or a production tail-latency guarantee.

Latency, Memory, and Energy Optimization

We now turn from correct model execution to three resource questions: response time, memory and energy. Their tradeoffs depend on the workload. Batching can change amortized compute cost while adding admission waiting; quantization can change storage and arithmetic while also changing predictions. Neither gives a fixed energy or latency factor without measurement. Choose settings against the deployment's actual quality and resource constraints, which differ between a data-center training job and a battery-powered robot.

Latency: Decomposition and Budgets

Choose the timing boundary before decomposing latency. For an illustrative serial control loop, from the start of sensor acquisition to completion of the commanded actuator response, use disjoint stages:

Tloop=Tcapture+Tqueue+Ttransport+Tpre+Tmodel+Tpost+Tencode+Tdecision+Tactuate\begin{aligned} T_{\text{loop}} ={}& T_{\text{capture}} + T_{\text{queue}} + T_{\text{transport}} + T_{\text{pre}} + T_{\text{model}} \\ &+ T_{\text{post}} + T_{\text{encode}} + T_{\text{decision}} + T_{\text{actuate}} \end{aligned}

where:

  • TcaptureT_{\text{capture}}: acquisition of the sensor data used for this decision.
  • TqueueT_{\text{queue}}: all waiting for admission, batching and resources along the measured path, excluding active work counted elsewhere.
  • TtransportT_{\text{transport}}: movement of inputs, outputs and commands between components, excluding their encoding and queue waiting.
  • TpreT_{\text{pre}}: input preparation such as resizing and normalization; learned encoder computation is counted under the model instead.
  • TmodelT_{\text{model}}: the model computation, including any learned observation encoder and dynamics calls used for this decision.
  • TpostT_{\text{post}}: output interpretation or uncertainty calculations not already in the model.
  • TencodeT_{\text{encode}}: serialization/deserialization of messages, not network transit.
  • TdecisionT_{\text{decision}}: remaining planner/controller work, excluding model calls already counted above.
  • TactuateT_{\text{actuate}}: command processing and actuator response within the chosen endpoint.

This sum assumes serial stages without overlap. Pipelined capture, asynchronous transfers or overlapping computation require critical-path timestamp accounting, not addition of every raw service duration. Count an operation once even if software names put it in two components. Optimizing model time still reduces the serial sum, but may yield only a small relative improvement when other terms dominate. The local benchmarks below cover a much narrower model/bookkeeping boundary, not this complete loop.

In a synchronous one-decision-per-period design at 30 Hz, the period is about 33.3 ms, and the scheduled work needs margin inside it. Asynchronous control may use different observation-age and stability requirements. A browser's display interval and input-to-display latency are distinct targets: 60 Hz means about 16.7 ms between frames, not a universal 100 ms frame budget. Offline imagination may prioritize throughput while still having completion deadlines, memory limits or feedback-staleness constraints. Choose the metric and its boundary for the actual consumer.

Concrete techniques:

  • Precision strategy. A tensor stored in fp16 or bf16 uses half the scalar bytes of fp32; that does not halve the entire footprint or guarantee a throughput factor. Long rollouts can amplify numerical perturbations in sensitive transitions, while contractions can attenuate them. Evaluate precision at one-step, rollout and decision levels on the target hardware. fp32 accumulation or a higher-precision reference may help diagnose an error, but periodic reference recomputation has its own cost and is not a generic cure. Quantization sensitivity depends on the operation and dynamics, not a universal transformer-versus-recurrence ordering.
  • Kernel fusion. Compatible fused kernels can reduce launches and intermediate memory traffic while preserving the intended computation. Fusing execution is not the same as algebraically folding input-dependent normalization into a fixed linear weight matrix. Measure launch and memory costs on the target runtime, and validate numerical differences introduced by a replacement kernel.
  • Sequence-length reduction. Dense full-prefix attention has quadratic arithmetic in context length; a cached one-token query still performs work linear in retained context, as discussed earlier. A fixed window or fixed-size recurrent state can bound per-step work at fixed model size. Chunking full-context attention bounds temporary working memory, not all pairwise arithmetic. Truncating context or changing architecture requires prediction and decision-quality validation.
  • Speculative decoding. Exact language-model sampling with a draft model requires a specified acceptance and correction procedure preserving the target distribution, as in Leviathan, Kalman and Matias. A fast approximate world-model transition checked against another model or incoming sensors is a different heuristic; it does not automatically inherit that exactness. Define the property being checked, the correction or fallback and the remaining approximation error. Bounded approximations can be intentional controller designs, but a timing gain alone does not establish their acceptability.

Memory: Anatomy of a Serving Footprint

Four useful tensor categories contribute memory pressure, before additional runtime overhead:

  • Weights. Parameters and any additional network copies actually loaded for inference. Their fraction of the footprint depends on model size and the other terms, not batch size alone.
  • Activations. Transient inference computation buffers; their sizes depend on the operation, batch and context. Ordinary inference does not retain backward activations.
  • Session state. Full-context KV storage grows with retained tokens and concurrent sessions. A fixed-width recurrent state grows with concurrent sessions but not with rollout duration unless extra history is retained.
  • Input/output buffers. These hold observations, actions, latent representations and predictions passed between encoders, dynamics models, decoders and policies. Multimodal inputs and outputs can make these buffers substantial.

Managing session state can be important, but profile the dominant term rather than assume it. For an explicitly disjoint tensor accounting, add weights, inference activations, retained session state and I/O, then account separately for runtime, allocator and workspace memory:

Mtotal=Mweights+Mactivations+Msessions+Mio+Moverhead\begin{aligned} M_{\text{total}} &= M_{\text{weights}} + M_{\text{activations}} + M_{\text{sessions}} + M_{\text{io}} + M_{\text{overhead}} \end{aligned}

where:

  • MweightsM_{\text{weights}}: bytes of loaded parameters and any required additional network copies
  • MactivationsM_{\text{activations}}: transient tensors for the current inference computation
  • MsessionsM_{\text{sessions}}: retained state for each session, with growth determined by its recurrent or attention representation
  • MioM_{\text{io}}: input and output buffers for encoders, decoders, and policies
  • MoverheadM_{\text{overhead}}: other allocated or reserved memory not assigned to the preceding terms, including runtime and workspaces

Strategies include:

  • Eviction policies. An evicted session can explicitly expire, with the next call returning an error, or it can be recovered by a defined protocol. Recovery may need a saved boundary state or observation context, subsequent actions, matching model/configuration and stochastic state; actions alone generally do not reconstruct it. LRU is only a selection policy, not a proof of recovery or request ordering. Silently replacing lost causal state with a default is the correctness hazard.
  • State compression. Quantization or a low-rank representation can reduce payload bytes, but is generally lossy; omitted information cannot simply be recomputed from the compressed state. Under a quantified contraction with the same future inputs, an initial-state perturbation can decay, but ongoing compression injects new error and control constraints may still be sensitive to it. Measure reconstruction, rollout and decision effects before accepting the memory tradeoff.
  • Sliding-window attention with action replay. When old observations become irrelevant, drop their KV entries but track the drop in a per-session ledger. Any subsequent request that references the dropped range is refused or reconstructed. This makes the memory/latency tradeoff explicit and auditable.

Length-aware grouping can change padding and allocation patterns. In an allocator requiring contiguous blocks, unusable free fragments can prevent an allocation even when total free bytes would suffice. Paged caches, preallocation and allocator policy can change that behavior. Monitor allocated, reserved and usable memory under changing session lengths; this CPU tutorial neither measures GPU fragmentation nor predicts an OOM time.

Energy: The Third Axis

Energy needs its own measurement. Battery capacity limits total energy, while electrical and thermal constraints limit power; data centers can have hard power and cooling limits as well as cost budgets. For a declared measurement boundary, energy is E=∫P(t) dtE=\int P(t)\,dt, where power PP is in watts and elapsed time is in seconds, giving joules. Runtime alone does not determine it. Account for the actual model lifecycle:

  • Training and adaptation may repeat. Model updates, imagined transitions and deployment queries all contribute; which dominates depends on their counts and measured costs, not whether the learner is called PPO or imagination-based.
  • Rollout length multiplies serial work. At a constant 10 ms per step, 10,000 sequential steps take 100 seconds before other overhead. Both horizon and per-step implementation are design choices to evaluate against decision quality. Time is not yet energy without power measurement.
  • Batching can amortize fixed overhead, but can also change active power, memory, waiting and unused work. Establish joules per useful inference from measured power and time, and test mobile deadlines and thermal limits rather than assume batching is unambiguously beneficial.

The common optimizations are:

  • Precision choices, assessed for numerical quality and measured energy on the target hardware.
  • Kernel fusion and reduced memory traffic, assessed with power instrumentation rather than treating bytes moved as an energy meter.
  • Avoiding redundant evaluations. Sampling-based planners may score many candidates. Valid repeated-prefix reuse can avoid computation when lookup and storage costs are lower; independent stochastic samples must retain their distinct sample identities. Pruning can reduce work but also change planning quality.

Energy is observable during development using a suitable meter or validated device counters and a declared scope. Measure idle and active intervals consistently, and include the components relevant to the decision, not only a convenient accelerator kernel. Energy per inference and per useful action answer different questions from average power; record the metrics the deployment budget actually constrains.

A Small Measurements Example

We measure a small CPU workload at several vectorized batch sizes. Each row times a different number of independent states advanced by one forward pass. Dividing that time by batch size gives amortized compute time per state, not end-to-end request latency. These timings vary with the CPU, runtime, thread settings and system load; they are not GPU throughput or energy measurements. Energy requires measured power integrated over time, as in the MLCommons power measurement rules.

In[15]:
Code
@torch.inference_mode()
def measure_batch(batch_size, steps=40, latent_dim=16, action_dim=2):
    z = torch.zeros(batch_size, latent_dim)
    a = torch.zeros(batch_size, action_dim)
    times = []
    for _ in range(3):
        model(z, a)  # warm-up
    for _ in range(steps):
        t0 = time.perf_counter()
        z = model(z, a)
        times.append(time.perf_counter() - t0)
    times = np.array(times)
    return float(np.median(times))


batch_sizes = [1, 2, 4, 8, 16, 32, 64, 128]
latencies = np.array([measure_batch(bs) for bs in batch_sizes])
per_session = latencies / batch_sizes
Out[16]:
Console
batch |  step ms |  per-session ms
    1 |    0.015 |          0.0148
    2 |    0.015 |          0.0074
    4 |    0.016 |          0.0039
    8 |    0.019 |          0.0024
   16 |    0.020 |          0.0012
   32 |    0.023 |          0.0007
   64 |    0.024 |          0.0004
  128 |    0.029 |          0.0002

Compare the actual table: batching may reduce amortized compute time even when a whole batch takes longer. The experiment does not simulate an arrival queue or establish a universal optimal batch size. A deployment must jointly measure throughput and end-to-end latency at its own request sizes and load.

Sequential Scaling Across Rollout Horizon

The next experiment isolates history-buffer bookkeeping during a recursive rollout. Both models use only the current state and action, so trimming the stored history cannot change predictions here. A transformer that attends to its history is different: truncating its cache may change the prediction distribution.

In[17]:
Code
@torch.inference_mode()
def measure_recursive(steps=400, latent_dim=16, action_dim=2):
    z = torch.zeros(1, latent_dim)
    a = torch.zeros(1, action_dim)
    history = torch.zeros(0, latent_dim)
    step_times = []
    for _ in range(steps):
        t0 = time.perf_counter()
        # naive recursive core
        z = model(z, a)
        history = torch.cat([history, z], dim=0)
        step_times.append(time.perf_counter() - t0)
    return np.array(step_times)


recursive_times = measure_recursive(steps=400)
print(
    f"Growing-history median per-step time: {np.median(recursive_times) * 1e3:.3f} ms"
)

Now the same rollout with a rolling history capped at a small window, so the working set stays bounded.

In[18]:
Code
@torch.inference_mode()
def measure_windowed(steps=400, latent_dim=16, action_dim=2, cap=32):
    if (
        isinstance(cap, bool)
        or not isinstance(cap, (int, np.integer))
        or cap <= 0
    ):
        raise ValueError("cap must be a positive integer")
    z = torch.zeros(1, latent_dim)
    a = torch.zeros(1, action_dim)
    history = torch.zeros(0, latent_dim)
    step_times = []
    for _ in range(steps):
        t0 = time.perf_counter()
        z = model(z, a)
        history = torch.cat([history, z], dim=0)
        if history.shape[0] > cap:
            history = history[-cap:].clone()
        step_times.append(time.perf_counter() - t0)
    return np.array(step_times)


windowed_times = measure_windowed(steps=400, cap=32)
print(
    f"Windowed-history median per-step time: {np.median(windowed_times) * 1e3:.3f} ms"
)

With fixed state dimension, repeatedly concatenating a growing history copies a number of elements proportional to ∑t=1Tt\sum_{t=1}^{T}t, giving quadratic total buffer-copy work. A fixed window of width WW bounds each copy, giving O(TW)O(TW) copy work, linear in horizon when WW is fixed. These are bookkeeping operation counts, not a guarantee that noisy CPU timings increase at every step. A preallocated buffer can avoid repeated copies without truncating history. Neither implementation measures attention or demonstrates a KV-cache speedup.

Out[19]:
Visualization
CPU step timings for growing and 32-state-capped history buffers.
Measured CPU time for a recursive model plus history-buffer concatenation. One buffer grows, the other is capped at 32 states. At this small scale overhead and noise may dominate; the operation-count analysis concerns copies, not attention or a guaranteed timing trend.

Read the observed variability rather than assuming either curve rises or stays flat. The analytical copy-count difference can exist even when tiny noisy timing measurements do not reveal it. Long-context attention has its own computation and memory costs, which must be measured separately.

The batch-size plot presents the amortized forward timings recorded earlier without assuming a knee occurs in the tested range.

Out[20]:
Visualization
Measured amortized CPU forward time per state versus batch size.
Amortized CPU forward time per state for the tested vectorized batch sizes. This excludes queue waiting and network costs and does not establish an end-to-end latency optimum or a universal batch-size knee.

Use this plot as a local compute diagnostic, not a serving configuration prescription. A production batching policy adds queue waiting, arrival rates, padding, memory limits and deadlines to the measured model cost. There need not be a single visible knee, and per-request latency does not universally grow linearly with batch size. Triton's batching documentation makes the deliberate queue-delay tradeoff explicit.

Limitations and Impact

The techniques in this chapter are powerful, and they also have hard limits. If a fixed model class cannot represent the needed dynamics, distributing its training does not remove that structural limitation. This differs from inaccurate fitted parameters that further training can improve. A cache that serves the wrong episode's state can also produce fast but incorrect rollouts. Systems infrastructure enables hypotheses to be tested; it does not validate them. A frequently overlooked consequence is that systems bugs tend to hide in the same corridors that reliability metrics use. If rollout latency grows when the queue is full, a poorly instrumented system may interpret the slower response as a longer-horizon measurement, not a delay. If a cache returns a stale latent, the rollout's apparent smoothness may look better than the actual dynamics. Tests must therefore be structured to probe the systems layer as explicitly as they probe the model. In particular, a well-designed system distinguishes model failures from infrastructure failures, and only then lets the evaluation ladder from The World-Model Evaluation Ladder be applied.

Distributed training introduces its own hazards. Incorrect loss normalization for unequal worker batches, mismatched optimizer state or incompatible numerical settings can compromise updates. Independent local sampling seeds alone do not make correctly synchronized replicas diverge. A slower rank can delay exposed synchronization; its percentage slowdown is not necessarily the whole job's percentage slowdown because overlap and other work matter. Per-rank timing, valid-token counts, losses and gradient norms make these effects observable. An aggregate loss can hide one faulty contribution, so inspect disaggregated evidence rather than assume every rank is healthy.

Checkpoint bugs include fresh optimizer moments on resume, repeated or skipped windows and missing auxiliary learning state. Define the state and test uninterrupted training against checkpointed continuation, comparing the next losses, parameter tensors and relevant optimizer or sampler state under a fixed deterministic configuration. Passing that test covers the exercised configuration only. Distributed shards, interruptions during writes, changed worker counts and device RNG state need additional tests. A mismatch calls for diagnosis: an incomplete snapshot and allowed numerical nondeterminism are different causes.

Three serving failures deserve explicit tests. First, non-preemptive batch work can block interactive requests sharing capacity. Protect their deadlines with admission control and an appropriate scheduling or reservation policy; separate queues are one implementation, not the only one. Second, a changed model, configuration or sampling policy must not reuse obsolete cache identities. Versioned keys can make old entries unreachable without physically deleting every response, while fresh-sample requests must not accidentally reuse a previous draw. Third, retained session state needs an owner and release policy: explicit close, expiry or eviction can prevent orphaned state from accumulating without bound. Test these lifecycle transitions and their errors as well as successful forward calls.

Measure energy per useful action as well as model throughput. If a planner evaluates thousands of candidates to choose one action, its decision energy includes every executed evaluation, not just the selected rollout. Valid prefix reuse or tested pruning can reduce discarded computation after accounting for overhead and planning-quality changes. This connects systems measurements to the algorithmic budgets discussed in Sampling-Based Planning and Model Predictive Control and Policies Learned in Imagination.

Finally, the biggest limitation is that systems quality is not transferable across model families without care. A serving stack optimized for a transformer world model with KV caches does not translate directly to a recurrent world model with hidden state. Variable-length episodes can supply fixed-length training windows. A training stack using those windows still needs correct temporal associations and episode-boundary handling; padding with appropriate masks is another design option. The general principles hold (distributed, cached, batched, latency-tuned), but the recipes are family-specific and frequently subtle. What works today for one flagship architecture is a starting point, not a finished design, for the next one.

The systems layer also supplies evidence for governance: dataset provenance, access logs, checkpoint retention and reproducible experiment records help an organization audit what it trained and deployed. Specific legal obligations depend on jurisdiction and use, and are outside this chapter's systems tutorial. Security, Ethics, Privacy, and Governance covers those governance concerns. The planned next chapter, Sim-to-Real Deployment and Continual Correction, considers noisy sensors, feedback delays and deployment shifts; it is not yet published.

Summary

Training and inference systems turn a world model from a mathematical object into something that behaves as a system under real constraints. The key ideas in this chapter:

  • Distributed-training axes include data (batch), tensor (layer operations), pipeline (depth) and context partitioning. Megatron's selected-activation sequence parallelism is distinct from context parallelism. Large models can combine axes; choose the mixture from memory constraints, interconnect and the actual workload.
  • Temporal partitioning is architecture-specific, not unique to world models. Ring exchange distributes attention over known context. Efficient recurrent scans require suitable compact associative transition structure; an arbitrary nonlinear recurrence still needs valid boundary states. Batching independent trajectories, checkpointing, recomputation or explicitly approximate training remain alternatives.
  • Data, checkpoint and experiment records support reproducible continuation and interpretable comparisons. Capture the algorithm's actual parameters, optimizer, sampler, random and auxiliary state; an EMA is required only if the algorithm uses it. Distinguish atomic publication from durable storage, and test recovery under the intended configuration.
  • Rollout serving carries causal and stochastic state. KV caches, session lifetimes and cache-aware batching need explicit ownership. Result-cache identity includes model version, initial state, action prefix, configuration and any reproducible RNG stream identity; fresh independent samples must not accidentally reuse a cached draw.
  • Latency, memory and energy require distinct measurement boundaries. A serial control-loop budget includes capture, queues, transport, input preparation, model work, output handling, encoding, decision and actuation; overlapping paths need critical-path accounting. Memory includes tensor categories plus runtime overhead, and fixed recurrent state differs from growing KV history. Energy per inference, per useful action and average power answer different budget questions.
  • Measure throughput and end-to-end latency together. A local batch microbenchmark need not show a knee and does not measure queue delay or energy. Choose settings against actual workload, memory and deadline constraints.

The systems layer supplies tests and records needed to assess reliability; it does not by itself certify prediction quality, safety or hard deadlines. The next chapter's deployment questions require checking these mechanisms against actual sensors, delays and shifts, not only the controlled CPU examples here.

Quiz

Ready to test your understanding? Take this quick quiz to reinforce what you've learned about training and inference systems for world models.

Training and Inference Systems Quiz

Question 1 of 80 of 8 completed
Why does the chapter separate memory capacity from throughput and wall-clock time as reasons for distributed training?

Comments

No comments yet. Be the first to share your thoughts!

Reference

Citation details

Cite or share this article.

BIBTEXAcademic
@misc{brenndoerfer2026traininginference, author = {Michael Brenndoerfer}, title = {Training and Inference Systems}, year = {2026}, url = {https://mbrenndoerfer.com/writing/training-and-inference-systems-world-model-serving}, organization = {mbrenndoerfer.com}, note = {Accessed: 2026-10-11} }
APAAcademic
Michael Brenndoerfer (2026). Training and Inference Systems. Retrieved from https://mbrenndoerfer.com/writing/training-and-inference-systems-world-model-serving
MLAAcademic
Michael Brenndoerfer. "Training and Inference Systems." 2026. Web. October 11, 2026. <https://mbrenndoerfer.com/writing/training-and-inference-systems-world-model-serving>.
CHICAGOAcademic
Michael Brenndoerfer. "Training and Inference Systems." Accessed October 11, 2026. https://mbrenndoerfer.com/writing/training-and-inference-systems-world-model-serving.
HARVARDAcademic
Michael Brenndoerfer (2026) 'Training and Inference Systems'. Available at: https://mbrenndoerfer.com/writing/training-and-inference-systems-world-model-serving (Accessed: October 11, 2026).
SimpleBasic
Michael Brenndoerfer (2026). Training and Inference Systems. https://mbrenndoerfer.com/writing/training-and-inference-systems-world-model-serving

About the author

Continue with the full handbook

This chapter is part of World Models Handbook. Use the handbook page to browse the complete table of contents and continue reading in sequence.

Explore World Models Handbook
Newsletter

Stay up to date

Get articles, book updates, and news delivered to your inbox.

No spam, unsubscribe anytime.

or

Join the community

Sign in to remove popups, track your reading progress, and join the discussion.