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.

Defining Models

A Model ties together regimes, an age grid, and a regime ID class into a solvable lifecycle model.

The Model Constructor

from lcm import Model

model = Model(
    regimes=regimes,  # dict mapping names to Regime instances
    ages=ages,  # AgeGrid defining the lifecycle timeline
    regime_id_class=RegimeId,  # @categorical dataclass mapping names to ScalarInt indices
    enable_jit=True,  # controls JAX compilation (default: True)
    fixed_params={},  # optional params baked in at init time
    description="",  # optional description string
)

All arguments are keyword-only. The three required arguments are regimes, ages, and regime_id_class. The finalized regimes are stored as model.user_regimes (plain Regime instances in user vocabulary); the processed canonical form is the engine-internal model._regimes.

Model-Level Regime Slots

When several regimes share functions, states, or actions, declare the shared structure once at the model level instead of repeating it per regime — a lifecycle model with a couple of dozen shared functions and a handful of shared states shrinks to one declaration site:

model = Model(
    regimes={"working": working, "retired": retired, "dead": dead},
    ages=ages,
    regime_id_class=RegimeId,
    functions={"taxes": taxes, "net_income": net_income},
    constraints={"budget": budget_constraint},
    states={"wealth": LinSpacedGrid(start=1, stop=100, n_points=50)},
    state_transitions={"wealth": next_wealth},
    actions={"consumption": LinSpacedGrid(start=1, stop=50, n_points=30)},
)

Each model-level slot accepts exactly what the regime-level slot accepts — including Phased, stochastic processes, per-target dicts, and fixed_transition. The entries are merged into every regime under three rules:

Pruning means a model-level state costs nothing in regimes that never touch it — the grid axis simply does not appear there. Two restrictions keep the device layout coherent: distributed=True (sharding) is legal only on model-level states, and a sharded state pruned from a non-terminal regime is an error (unshard it or make the regime use it).

Regime ID Classes

The regime_id_class maps regime names to integer indices. Use the @categorical decorator to create it:

from lcm import categorical
from lcm.typing import ScalarInt


@categorical(ordered=False)
class RegimeId:
    retired: ScalarInt
    working: ScalarInt

Rules:

Age Grids

The ages argument defines the lifecycle timeline. There are two construction modes:

Range-based

from lcm import AgeGrid

ages = AgeGrid(start=25, stop=75, step="Y")  # annual steps, ages 25 to 75

Step formats:

The stop value is inclusive if (stop - start) is exactly divisible by the step size.

Exact values

ages = AgeGrid(exact_values=[25, 35, 45, 55, 65, 75])

Use this for irregular age spacing.

Key properties

Model Validation Rules

The Model constructor validates:

Inspecting a Model

After construction, the model exposes several useful attributes:

model.user_regimes  # immutable mapping of finalized `Regime` objects
model.pruned_variables  # per regime, the broadcast names pruned by DAG reachability
model.n_periods  # number of periods
model.regime_names_to_ids  # name -> integer mapping
model.get_params_template()  # mutable copy of the parameter template

Use model.get_params_template() to get a mutable copy of the parameter template — see Parameters.

Complete Example

import jax.numpy as jnp
from lcm import AgeGrid, DiscreteGrid, LinSpacedGrid, Model, Regime, categorical
from lcm.typing import ScalarInt


@categorical(ordered=False)
class RegimeId:
    retired: ScalarInt
    working: ScalarInt


@categorical(ordered=True)
class LaborSupply:
    do_not_work: ScalarInt
    work: ScalarInt


def next_wealth(wealth, consumption, interest_rate):
    return (wealth - consumption) * (1 + interest_rate)


def next_regime(labor_supply):
    return jnp.where(
        labor_supply == LaborSupply.work, RegimeId.working, RegimeId.retired
    )


def utility(consumption, labor_supply, disutility_of_work):
    return jnp.log(consumption) - disutility_of_work * labor_supply


def terminal_utility(wealth):
    return jnp.log(wealth)


working = Regime(
    transition=next_regime,
    states={
        "wealth": LinSpacedGrid(start=1, stop=100, n_points=50),
    },
    state_transitions={
        "wealth": next_wealth,
    },
    actions={
        "consumption": LinSpacedGrid(start=1, stop=50, n_points=30),
        "labor_supply": DiscreteGrid(LaborSupply),
    },
    functions={"utility": utility},
)

retired = Regime(
    transition=None,
    states={
        "wealth": LinSpacedGrid(start=1, stop=100, n_points=50),
    },
    functions={"utility": terminal_utility},
)

model = Model(
    regimes={"working": working, "retired": retired},
    ages=AgeGrid(start=25, stop=75, step="Y"),
    regime_id_class=RegimeId,
)

Correlated State Transitions

Use MarkovTransition when one state has its own stochastic law. When several target states must use the same realization, declare one edge-owned JointTransition instead. The outer key is the target regime; the inner key is the transition-local node name supplied to every output law.

import jax.numpy as jnp

from lcm import JointTransition, MarkovTransition, Regime, fixed_transition
from lcm.typing import FloatND, IntND


MATCH_SUPPORT = {
    "partner_wealth": jnp.asarray([0.5, 2.0]),
    "partner_health": jnp.asarray([0, 1], dtype=jnp.int32),
}


def match_probabilities(match_type: IntND, match_weights: FloatND) -> FloatND:
    return match_weights[match_type]


def next_wealth(wealth, partner_match):
    return wealth + partner_match["partner_wealth"]


def next_partner_health(partner_match):
    return partner_match["partner_health"].astype(jnp.int32)


single = Regime(
    transition={"couple": MarkovTransition(probability_of_couple)},
    states={"wealth": wealth_grid, "match_type": match_type_grid},
    state_transitions={"match_type": fixed_transition("match_type")},
    joint_transitions={
        "couple": {
            "partner_match": JointTransition(
                support_size=2,
                support=MATCH_SUPPORT,
                probabilities=match_probabilities,
                outputs={
                    "wealth": next_wealth,
                    "partner_health": next_partner_health,
                },
            ),
        },
    },
    actions=single_actions,
    functions={"utility": single_utility},
)

Every support leaf has leading axis support_size; one aligned row is enumerated in solve and sampled in simulation. The latent partner_match is not a state or grid, adds no value-function axis, and never appears in initial conditions or simulation results. The expectation over its nodes is formed inside the action value before the source action is maximized.

A callable support may read only period, age, and parameters. Probabilities may also read source states, actions, and helpers. Output laws may transform the shared node using source values and may read already-resolved next_<state> outputs on the same target edge. Invalid support shapes and probabilities are rejected during the params-bound runtime preflight; probabilities are never silently normalized.

An output may target a stochastic process as well as an ordinary grid, which is how correlated innovations land on a grid pylcm discretized for you rather than one you discretized by hand. The output law still names a physical value; because the target’s value function is stored on the process’s nodes, that value reaches the continuation as its coefficients in the node basis — the hat weights of linear interpolation. Naming a node reads that node alone; naming a point between nodes reads the linear interpolation of the target’s value function, which is the only reading its nodes support. The output law displaces the process’s own law on that edge, so the correlation the kernel imposes is what the target is entered at. The support is the contract: a value outside the process’s grid has no representation in that basis and yields NaN, which the caller’s value function reports rather than extrapolating.

Parameters for support and probabilities live below the kernel name, while output parameters keep the ordinary target-local next_<state> paths:

params[source][target][kernel]["support"]
params[source][target][kernel]["probabilities"]
params[source][target]["next_wealth"]

Use Phased(solve=JointTransition(...), simulate=JointTransition(...)) around the whole kernel for perceived versus realized dynamics. Both variants must retain the same output names, support size, and literal support schema.

See Also