Skip to article frontmatterSkip to article content
Site not loading correctly?

This may be due to an incorrect BASE_URL configuration. See the MyST Documentation for reference.

Performance and memory tuning

Tune a measured model, not an abstract solver. Correctness comes first: choose a solver whose assumptions represent the economics, run the model at log_level="debug", and compare a reduced problem with grid search before optimizing it.

The performance workflow has four measurements:

  1. cold environment and cold compilation cache;

  2. cold model compilation in an installed environment;

  3. warm execution in the same process;

  4. peak host and device memory.

Record model size, precision, device, solver configuration, and whether the compilation cache was warm. Without those fields, two timings are not comparable.

First locate the limiting axis

List the sizes of:

The scaling discussion explains how those axes enter each solver family. The largest declared grid is not necessarily the largest intermediate: a product or envelope matrix can dominate.

Reduce discretization only when accuracy permits

Fewer nodes reduce work and memory, but the right grid is an economic approximation choice. Inspect value and policy changes as grids are refined. Put resolution near curvature, boundaries, or regions visited frequently in simulation.

PiecewiseLinSpacedGrid and PiecewiseLogSpacedGrid control density around known locations. They do not declare a budget kink or cliff to NBEGM; use the structured budget declarations for that.

Stream work with explicit batch widths

Some controls reduce live intermediates. Grid batch_size, stochastic_node_batch_size, envelope_segment_block_size, subject_batch_size, and any solver field whose Reference contract explicitly says it streams an evaluation axis can lower temporary workspace. The exact effect still depends on retained banks and downstream folds; for example, NEGM.outer_batch_size can lower temporary evaluation memory without capping the retained candidate bank.

NBEGM’s interval_batch_size, cell_block_size, and branch_batch_size are compiled batch widths for the corresponding lax.map axes. A positive value smaller than the axis bounds how many entries are evaluated together; 0, or a value covering the axis, uses one vectorized pass. Lower values can reduce live intermediates inside that mapped core at the cost of more sequential execution. They do not cap surrounding arrays, retained candidate banks, compilation memory, or total device memory.

Choose the largest batch that meets the measured memory target, then verify values and runtime against the whole-axis setting on the model and backend you will use.

Exact solver fields are in Solvers and capabilities, Upper envelopes, and Outer search.

Distribute independent discrete state work

distributed=True shards a supported discrete grid over visible devices. Continuous grids reject distribution because their interpolation needs the full coordinate axis. A grid cannot be both batched and distributed; if a shard remains too large, batch a different axis.

Before solving, verify the resources actually visible to JAX:

import jax

assert jax.device_count() == expected_devices

A larger GPU can run larger chunks and may benefit from more concurrent independent work. It does not automatically shorten a workload made of small sequential kernels. Measure occupancy and memory rather than extrapolating from device memory alone.

Batch forward simulation

model.simulate(subject_batch_size=k, ...) bounds the subject workspace and offloads completed chunks to host. Random keys are assigned by global subject index, so changing the batch size does not change simulated draws.

The same fixed seed also gives the same EV1 taste-shock choices in lazy and ahead-of-time simulation. Subject chunking and Model(n_subjects=...) change compilation and workspace shape, not which per-subject Gumbel key is used. Keep the seed, parameters, initial conditions, and model fixed when using that invariance as a regression check.

If Model(n_subjects=n) was constructed, a matching first simulation can compile for that population/chunk shape ahead of execution and cache it. Reuse requires stable parameter shapes and dtypes.

Reuse compilation

pylcm enables a persistent JAX compilation cache by default. Check:

import jax

print(jax.config.jax_compilation_cache_dir)

A None value means no persistent cache. JAX_COMPILATION_CACHE_DIR chooses the full path; LCM_COMPILATION_CACHE_NAME chooses a project leaf under the default root. Compilation keys still change when program shapes or the software environment change.

Do not create JAX device arrays at module import solely for constants. Keep tables as NumPy arrays and convert inside traced functions; this avoids initializing a device in processes that only import the model.

Runtime environment controls are listed in Runtime, results, and persistence.

Benchmark the decision you face

For a solver comparison, hold the economic model and accuracy target fixed. Report:

Use the external LCM solver benchmarks for evolving shared evidence. The pylcm package’s own regression benchmarks belong to the Development chapter.

Checklist