Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
18 commits
Select commit Hold shift + click to select a range
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 4 additions & 4 deletions .github/workflows/main.yml
Original file line number Diff line number Diff line change
Expand Up @@ -30,12 +30,12 @@ jobs:
- uses: actions/setup-python@v6
with:
python-version: ${{ matrix.python-version }}
- name: Install pylcm from feat/nb-egm
# This slice depends on the NB-EGM solver, which lives on pylcm's
# feat/nb-egm branch; revert to @main once that branch merges.
- name: Install pylcm at the revision aca-dev pins
# Pinned by commit so a deleted or rewritten branch cannot break the
# install; move it together with the aca-dev pylcm submodule pointer.
run: >-
pip install "pylcm @
git+https://github.com/OpenSourceEconomics/pylcm.git@feat/nb-egm"
git+https://github.com/OpenSourceEconomics/pylcm.git@4f73efe34fd929b119cd10c9722bea5c10c2caf2"
- name: Install aca-model with test deps
run: pip install -e . pytest pdbp
- name: Run pytest
Expand Down
40 changes: 40 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
@@ -1,3 +1,43 @@
# aca-model

Core lifecycle model for the ACA structural retirement project.

Model factories accept `execution_config=ExecutionConfig(...)`. An omitted policy uses
the smallest JAX allocator `bytes_limit` among the selected accelerators as an explicit
per-device budget. This enables budget-aware width selection; CPU construction remains
unbudgeted. Missing accelerator limits require an explicit policy instead of silently
falling back to unbudgeted execution.

To combine a measured budget with device selection and sharding:

```python
from dataclasses import replace

from aca_model.execution import execution_config_for_devices

execution_config = replace(
execution_config_for_devices(devices=(0, 1, 2)),
sharded_states=("pref_type",),
)
```

Pass this policy to the baseline, ACA or benchmark factory. Record its exact
`device_memory_bytes`, the selected devices' allocator observations, executed precision
and selected program widths with performance evidence. The budget is an allocator
ceiling; it does not establish that all workload phases fit.

Explicit policies pass through unchanged. In particular, `ExecutionConfig()` is an
unbudgeted control and retains pylcm's bootstrap width selection. Economic grids,
numerical solver settings and result retention are separate choices.

`GridConfig` accepts economic grid sizes and numerical solver choices only. Execution
controls belong to `ExecutionConfig`: sharded state names go in `sharded_states`,
declared program widths in `axis_widths`, and the per-device byte ceiling in
`device_memory_bytes`. Obsolete grid execution keywords are rejected. A state-specific
batch size has no automatic conversion to a flattened program width; callers must choose
and record an explicit execution policy.

For comparisons with an older pylcm version, retain its model/configuration pin and use
a separately reviewed API adapter. Preserve grids, transition declarations, numerical
solver choices, retention and seeds across the comparison. An unspecified width in the
new policy means planner selection, not a translation of an old zero-valued batch size.
31 changes: 31 additions & 0 deletions docs/dag-migration.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,31 @@
# ACA: one structural stage table, existing numerical programs

Requires the proposed pylcm v3 schedule/support API, not unchanged PR #474.
Both factories forward optional `initial_regimes` unchanged. Default coverage
requires no separate population/activity table; an explicit entry restriction
changes admission, never solved values. Coverage remains exactly **182** nodes.

`_AGE_STAGES` and the host-side clock supply both numerical next-age routing and
source-age support schedules. Each source's existing numerical transition body and
per-target probability cell are shared across age cases. There is no closure
specialization per boundary. Code order/global regime IDs and existing parameter
paths remain; different target schemas can still need different continuation
programs, so no backend compile-time improvement is claimed.

The dead regime remains `regime_transitions=None` at every age, as in the supplied base.
Age 94 retains its living consumption/saving problem and handoff into the separate
age-95 bequest. The three-to-two health grid, claim/lagged-work entry and exit,
carried pension state and phase semantics are not simplified away.

Use `model.reachability.nodes` and `model.initial_nodes`; keep the existing phase
target queries and period-keyed solutions. Normalize schedules before existing
solver validation. A wrapper alone does not make a supported solver case invalid;
genuine unsupported variation must raise via current errors, without silently
changing actions or solver. GridSearch is the first numerical acceptance target.

The original transition test module is retained byte-for-byte; new coverage tests
check nodes, boundary support, shared probability cells and terminal declarations.
Offline JAX row equality is not real model construction, Bellman/panel equivalence,
compiled validation/handoff/replay acceptance or GPU evidence. Run those gates and
matched build/compile/warm timings in the supported package environment. Annual
ACA calibration is not automatically valid on a differently spaced clock.
9 changes: 8 additions & 1 deletion src/aca_model/__init__.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,13 @@
import os

import jax

jax.config.update("jax_enable_x64", True)
_x64_requested = os.environ.get("ACA_JAX_ENABLE_X64", "1").lower() not in {
"0",
"false",
"no",
}
jax.config.update("jax_enable_x64", _x64_requested)

# Import lcm before installing the claw so its `_jaxtyping_patch` (picklable
# jaxtyping sentinel) and `MappingProxyType` pytree registration are in place.
Expand Down
30 changes: 19 additions & 11 deletions src/aca_model/aca/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,19 +7,19 @@
from collections.abc import Mapping
from typing import Any

from lcm import AgeGrid, DiscreteGrid, Model
from lcm.typing import UserParams
from lcm import AgeGrid, DiscreteGrid, ExecutionConfig, Model
from lcm.typing import InitialRegimes, UserParams

from aca_model.aca import PolicyVariant
from aca_model.aca.regimes import build_all_regimes
from aca_model.baseline.model import _fail_if_dcegm_without_consumption_points
from aca_model.baseline.regimes import RegimeId, SolverName, build_model_slots
from aca_model.config import MODEL_CONFIG, GridConfig
from aca_model.config import MODEL_AGES, GridConfig
from aca_model.execution import execution_config_for_devices


def create_model(
*,
n_subjects: int,
policy: PolicyVariant,
fixed_params: UserParams,
wage_params: Mapping[str, Any],
Expand All @@ -28,11 +28,12 @@ def create_model(
pref_type_grid: DiscreteGrid,
solver: SolverName = "brute_force",
consumption_dollars_points: tuple[float, ...] | None = None,
execution_config: ExecutionConfig | None = None,
initial_regimes: InitialRegimes | None = None,
) -> Model:
"""Create an ACA policy variant model.

Args:
n_subjects: Forwarded to `lcm.Model(n_subjects=...)`.
policy: Which ACA policy combination to apply (e.g.
`PolicyVariant.ACA`).
fixed_params: Parameters to fix at model creation time. Pass
Expand All @@ -53,16 +54,18 @@ def create_model(
consumption_dollars_points: Construction-time consumption action
gridpoints; required under DC-EGM. See
`aca_model.baseline.model.create_model`.
initial_regimes: Optional simulation-entry contract, not a solve-domain
restriction. None permits every covered node; an empty mapping makes
the model solve-only. Empirical starting pairs stay in the data/recipe.
execution_config: Explicit hardware-local policy forwarded unchanged.
None uses the smallest selected accelerator allocator limit as the
device-memory budget; CPU construction remains unbudgeted.

Returns:
pylcm Model.

"""
ages = AgeGrid(
start=MODEL_CONFIG.start_age,
stop=MODEL_CONFIG.end_age - 1,
step="Y",
)
ages = AgeGrid(exact_values=MODEL_AGES)
_fail_if_dcegm_without_consumption_points(
solver=solver, consumption_dollars_points=consumption_dollars_points
)
Expand All @@ -87,11 +90,16 @@ def create_model(

return Model(
regimes=regimes,
initial_regimes=initial_regimes,
ages=ages,
regime_id_class=RegimeId,
description=f"Structural retirement model ({policy.name})",
fixed_params=fixed_params,
derived_categoricals=derived_categoricals,
n_subjects=n_subjects,
execution_config=(
execution_config_for_devices()
if execution_config is None
else execution_config
),
**model_slots,
)
33 changes: 24 additions & 9 deletions src/aca_model/agent/preferences.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,13 @@
# endowment; it only bends the map near and beyond the endowment.
_LEISURE_SMOOTHING_FRACTION = 0.01

# Positive floor that leisure approaches as work costs exceed the endowment, as a
# fraction of the time endowment. It keeps felicity and its consumption derivative
# inside float32's range when leisure_available is far below zero, down to the
# NB-EGM inverse bracket's lowest consumption; away from the floor its effect decays
# exponentially.
_LEISURE_FLOOR_FRACTION = 1e-5


@categorical(ordered=False)
class PrefType:
Expand Down Expand Up @@ -71,17 +78,25 @@ def fixed_cost_of_work(
def _smooth_leisure_floor(
leisure_available: FloatND, time_endowment: ScalarFloat
) -> FloatND:
"""Bend leisure to a strictly positive floor as work costs approach the endowment.

`softplus(x) = log(1 + e^x)` via `jnp.logaddexp(0, x)`, scaled by a small fraction
of the endowment. Where `leisure_available` is large relative to the smoothing width
the map reduces to `leisure_available` (bulk unchanged); as it falls to zero leisure
bends to `0⁺` — never negative, never a kinked clamp — so the CRRA aggregator never
receives a non-positive base. The smoothing width scales with the endowment, so the
map is scale-invariant.
"""Bend leisure smoothly onto a positive floor as work costs approach the endowment.

Leisure is a smooth maximum of `leisure_available` and the floor
`F = _LEISURE_FLOOR_FRACTION * time_endowment`:
`s * logaddexp(F / s, leisure_available / s)`, with smoothing width
`s = _LEISURE_SMOOTHING_FRACTION * time_endowment`.

- `leisure_available` large relative to `s`: leisure equals `leisure_available`
up to a term that decays like `e^(-leisure_available / s)`.
- `leisure_available` at or below zero: leisure bends to `F⁺` — never below the
floor, never a kinked clamp — so every work choice stays feasible and the CRRA
aggregator's base stays large enough for float32 felicity and marginal
felicity.

Both widths scale with the endowment, so the map is scale-invariant.
"""
smoothing = _LEISURE_SMOOTHING_FRACTION * time_endowment
return smoothing * jnp.logaddexp(0.0, leisure_available / smoothing)
floor = _LEISURE_FLOOR_FRACTION * time_endowment
return smoothing * jnp.logaddexp(floor / smoothing, leisure_available / smoothing)


def leisure_canwork_retiree_or_nongroup(
Expand Down
32 changes: 20 additions & 12 deletions src/aca_model/baseline/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,41 +5,42 @@

Usage:
from aca_model.baseline.model import create_model
model = create_model(n_subjects=..., fixed_params=..., wage_params=..., ...)
model = create_model(fixed_params=..., wage_params=..., ...)
params = get_default_params()
V = model.solve(params)
"""

from collections.abc import Mapping
from typing import Any

from lcm import AgeGrid, DiscreteGrid, Model
from lcm.typing import UserParams
from lcm import AgeGrid, DiscreteGrid, ExecutionConfig, Model
from lcm.typing import InitialRegimes, UserParams

from aca_model.baseline.regimes import (
RegimeId,
SolverName,
build_all_regimes,
build_model_slots,
)
from aca_model.config import MODEL_CONFIG, GridConfig
from aca_model.config import MODEL_AGES, GridConfig
from aca_model.execution import execution_config_for_devices


def create_model(
*,
n_subjects: int,
fixed_params: UserParams,
wage_params: Mapping[str, Any],
derived_categoricals: Mapping[str, DiscreteGrid],
grid_config: GridConfig,
pref_type_grid: DiscreteGrid,
solver: SolverName = "brute_force",
consumption_dollars_points: tuple[float, ...] | None = None,
execution_config: ExecutionConfig | None = None,
initial_regimes: InitialRegimes | None = None,
) -> Model:
"""Create the baseline structural retirement model.

Args:
n_subjects: Forwarded to `lcm.Model(n_subjects=...)`.
fixed_params: Parameters to fix at model creation time. Fixed
params are partialled into compiled functions and removed
from the params template. Pass data-derived constants here;
Expand Down Expand Up @@ -69,17 +70,19 @@ def create_model(
continuous-action grid at model construction); `None` keeps
the runtime-points grid completed per iteration via
`inject_consumption_dollars_points`.
initial_regimes: Optional simulation-entry contract, not a solve-domain
restriction. None permits every covered node; an empty mapping makes
the model solve-only. Empirical starting pairs stay in the data/recipe.
execution_config: Explicit hardware-local policy forwarded unchanged.
None uses the smallest selected accelerator allocator limit as the
device-memory budget; CPU construction remains unbudgeted.

Returns:
A pylcm Model with 19 regimes (18 non-terminal + dead) spanning
ages 51-95. Regime names follow the `<his>_<medicare>_<ss>_<work>` scheme.

"""
ages = AgeGrid(
start=MODEL_CONFIG.start_age,
stop=MODEL_CONFIG.end_age - 1,
step="Y",
)
ages = AgeGrid(exact_values=MODEL_AGES)
_fail_if_dcegm_without_consumption_points(
solver=solver, consumption_dollars_points=consumption_dollars_points
)
Expand All @@ -101,12 +104,17 @@ def create_model(

return Model(
regimes=regimes,
initial_regimes=initial_regimes,
ages=ages,
regime_id_class=RegimeId,
description="Baseline structural retirement model (pre-ACA)",
fixed_params=fixed_params,
derived_categoricals=derived_categoricals,
n_subjects=n_subjects,
execution_config=(
execution_config_for_devices()
if execution_config is None
else execution_config
),
**model_slots,
)

Expand Down
11 changes: 3 additions & 8 deletions src/aca_model/baseline/regimes/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,11 +36,6 @@
from aca_model.baseline.regimes._nbegm import build_nbegm_solver
from aca_model.config import GridConfig

# NBEGM is a per-regime (not global) solver: it solves one 1-D
# consumption/savings regime with at most one discrete action, so it attaches
# only to the M1 vertical-slice regime.
_NBEGM_REGIME = "nongroup_nomc_inelig_canwork"

__all__ = [
"REGIME_SPECS",
"RegimeId",
Expand Down Expand Up @@ -115,7 +110,7 @@ def build_all_regimes(
name,
grids,
dcegm_solver=dcegm_solver,
nbegm_solver=nbegm_solver if name == _NBEGM_REGIME else None,
nbegm_solver=nbegm_solver,
)
regimes["dead"] = build_dead_regime(solver=solver)
return regimes
Expand All @@ -136,7 +131,7 @@ def build_model_slots(
Both the baseline and the ACA `create_model` consume this — the ACA
overlay swaps only regime-level functions, so the broadcast slots are
policy-invariant. Under DC-EGM the solver-contract functions join the
broadcast set and no borrowing constraint is declared.
broadcast set; the borrowing constraint is declared under every solver.
"""
grids = build_grids(
grid_config=grid_config,
Expand All @@ -146,7 +141,7 @@ def build_model_slots(
)
return {
"functions": build_model_functions(solver=solver),
"constraints": build_model_constraints(),
"constraints": build_model_constraints(solver=solver),
"states": build_model_states(grids),
"state_transitions": build_model_state_transitions(),
}
Loading
Loading